Compare commits
9 commits
main
...
test-v100-
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6eaa2dfc49 | ||
|
|
1307f24c80 | ||
|
|
e0becf489f | ||
|
|
ca826bc3d8 | ||
|
|
a1e2333d7c | ||
|
|
690185abe1 | ||
|
|
cca441f125 | ||
|
|
84f76a42cb | ||
|
|
8b1611c26c |
2 changed files with 279 additions and 11 deletions
182
tests/python/test_v100_fullft_precision.py
Normal file
182
tests/python/test_v100_fullft_precision.py
Normal file
|
|
@ -0,0 +1,182 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved.
|
||||
|
||||
"""Regression tests for full finetuning precision on no-bf16 GPUs (V100/T4).
|
||||
|
||||
Full finetuning upcasts trainable weights to float32, so the model dtype is
|
||||
float32 (not bfloat16). The SFTTrainer mixed-precision template in
|
||||
unsloth/models/rl.py must then:
|
||||
- run the forward pass under float16 autocast for normal models,
|
||||
- keep FORCE_FLOAT32 models (Gemma3, gpt_oss, ...) in pure float32,
|
||||
- never select bf16 on hardware without bf16.
|
||||
|
||||
We execute the REAL template block extracted from rl.py source (no heavy unsloth
|
||||
import) against mocked inputs. See issue #4082.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
torch = pytest.importorskip("torch")
|
||||
|
||||
RL_PY = Path(__file__).resolve().parents[2] / "unsloth" / "models" / "rl.py"
|
||||
|
||||
|
||||
def _extract_mixed_precision_code() -> str:
|
||||
lines = RL_PY.read_text().split("\n")
|
||||
try:
|
||||
start = next(i for i, l in enumerate(lines) if "mixed_precision = (" in l)
|
||||
except StopIteration:
|
||||
pytest.skip("mixed_precision template not found in rl.py")
|
||||
body, k = [], start + 1
|
||||
while lines[k].strip() != ")":
|
||||
body.append(lines[k])
|
||||
k += 1
|
||||
return eval("(\n" + "\n".join(body) + "\n)") # only string literals + comments
|
||||
|
||||
|
||||
CODE = _extract_mixed_precision_code()
|
||||
|
||||
|
||||
def _decide(
|
||||
dtype,
|
||||
*,
|
||||
bf16_supported,
|
||||
force_float32,
|
||||
full_finetuning,
|
||||
mixed_precision,
|
||||
fp16,
|
||||
bf16,
|
||||
):
|
||||
"""Run the template block; return (args.fp16, args.bf16, ACCELERATE_MP, raised)."""
|
||||
uz = types.ModuleType("unsloth_zoo")
|
||||
uzu = types.ModuleType("unsloth_zoo.utils")
|
||||
uzu._get_dtype = lambda x: x
|
||||
uzd = types.ModuleType("unsloth_zoo.device_type")
|
||||
uzd.device_is_bf16_supported = lambda: bf16_supported # device-aware signal stub
|
||||
sys.modules.setdefault("unsloth_zoo", uz)
|
||||
sys.modules["unsloth_zoo.utils"] = uzu
|
||||
sys.modules["unsloth_zoo.device_type"] = uzd
|
||||
for k in (
|
||||
"UNSLOTH_FORCE_FLOAT32",
|
||||
"UNSLOTH_ENABLE_FULL_FINETUNING",
|
||||
"UNSLOTH_MIXED_PRECISION",
|
||||
"ACCELERATE_MIXED_PRECISION",
|
||||
):
|
||||
os.environ.pop(k, None)
|
||||
os.environ["UNSLOTH_FORCE_FLOAT32"] = "1" if force_float32 else "0"
|
||||
os.environ["UNSLOTH_ENABLE_FULL_FINETUNING"] = "1" if full_finetuning else "0"
|
||||
os.environ["UNSLOTH_MIXED_PRECISION"] = mixed_precision
|
||||
orig = torch.cuda.is_bf16_supported
|
||||
torch.cuda.is_bf16_supported = lambda *a, **k: bf16_supported
|
||||
args = types.SimpleNamespace(fp16 = fp16, bf16 = bf16, mixed_precision = None)
|
||||
emb = types.SimpleNamespace(weight = types.SimpleNamespace(dtype = dtype))
|
||||
model = types.SimpleNamespace(
|
||||
config = types.SimpleNamespace(dtype = dtype, torch_dtype = dtype),
|
||||
get_input_embeddings = lambda: emb,
|
||||
)
|
||||
raised = None
|
||||
try:
|
||||
exec(CODE, {"torch": torch, "os": os}, {"args": args, "model": model})
|
||||
except TypeError:
|
||||
raised = "TypeError"
|
||||
finally:
|
||||
torch.cuda.is_bf16_supported = orig
|
||||
return args.fp16, args.bf16, os.environ.get("ACCELERATE_MIXED_PRECISION"), raised
|
||||
|
||||
|
||||
def test_v100_normal_fullft_fp16_explicit():
|
||||
# Normal model, full FT (weights upcast to float32), V100, fp16=True.
|
||||
fp16, bf16, amp, raised = _decide(
|
||||
torch.float32,
|
||||
bf16_supported = False,
|
||||
force_float32 = False,
|
||||
full_finetuning = True,
|
||||
mixed_precision = "float32",
|
||||
fp16 = True,
|
||||
bf16 = False,
|
||||
)
|
||||
assert raised is None
|
||||
assert (fp16, bf16) == (True, False) # float32 weights + fp16 forward
|
||||
|
||||
|
||||
def test_v100_normal_fullft_precision_unset():
|
||||
# Same, but user left precision unset -> must pick fp16, never bf16.
|
||||
fp16, bf16, amp, raised = _decide(
|
||||
torch.float32,
|
||||
bf16_supported = False,
|
||||
force_float32 = False,
|
||||
full_finetuning = True,
|
||||
mixed_precision = "float32",
|
||||
fp16 = False,
|
||||
bf16 = False,
|
||||
)
|
||||
assert raised is None
|
||||
assert (fp16, bf16) == (True, False)
|
||||
assert amp == "fp16"
|
||||
|
||||
|
||||
def test_force_float32_model_fullft_is_pure_float32():
|
||||
# FORCE_FLOAT32 model (Gemma3, gpt_oss, ...) in full FT -> pure float32, no autocast.
|
||||
fp16, bf16, amp, raised = _decide(
|
||||
torch.float32,
|
||||
bf16_supported = False,
|
||||
force_float32 = True,
|
||||
full_finetuning = True,
|
||||
mixed_precision = "float32",
|
||||
fp16 = True,
|
||||
bf16 = False,
|
||||
)
|
||||
assert raised is None
|
||||
assert (fp16, bf16) == (False, False)
|
||||
assert amp in (None, "no")
|
||||
|
||||
|
||||
def test_no_bf16_on_volta_in_auto_branch():
|
||||
# bf16 model dtype but no bf16 HW, precision unset -> fp16, never bf16.
|
||||
fp16, bf16, amp, raised = _decide(
|
||||
torch.bfloat16,
|
||||
bf16_supported = False,
|
||||
force_float32 = False,
|
||||
full_finetuning = False,
|
||||
mixed_precision = "float32",
|
||||
fp16 = False,
|
||||
bf16 = False,
|
||||
)
|
||||
assert bf16 is False
|
||||
|
||||
|
||||
def test_bf16_gpu_unchanged_auto_branch():
|
||||
# Regression guard: on a bf16 GPU, a float32 model with unset precision
|
||||
# still selects bf16 autocast (behavior must not change for bf16 hardware).
|
||||
fp16, bf16, amp, raised = _decide(
|
||||
torch.float32,
|
||||
bf16_supported = True,
|
||||
force_float32 = False,
|
||||
full_finetuning = True,
|
||||
mixed_precision = "float32",
|
||||
fp16 = False,
|
||||
bf16 = False,
|
||||
)
|
||||
assert raised is None
|
||||
assert (fp16, bf16) == (False, True)
|
||||
|
||||
|
||||
def test_genuine_bf16_model_with_fp16_still_raises():
|
||||
# A real bfloat16 model on bf16 HW with fp16 requested is a genuine mismatch.
|
||||
_, _, _, raised = _decide(
|
||||
torch.bfloat16,
|
||||
bf16_supported = True,
|
||||
force_float32 = False,
|
||||
full_finetuning = False,
|
||||
mixed_precision = "float32",
|
||||
fp16 = True,
|
||||
bf16 = False,
|
||||
)
|
||||
assert raised == "TypeError"
|
||||
|
|
@ -580,6 +580,31 @@ _UNSLOTH_GRPO_HIDDEN_STATES_WRAPPED_ATTR = "_unsloth_grpo_hidden_states_forward_
|
|||
_UNSLOTH_GRPO_HIDDEN_STATES_WARNING_ATTR = "_unsloth_grpo_hidden_states_warning_issued"
|
||||
|
||||
|
||||
def _grpo_owns_lm_head(module):
|
||||
# Does this module apply lm_head itself (i.e. its forward emits `.logits`)?
|
||||
if module is None:
|
||||
return False
|
||||
if getattr(module, "lm_head", None) is not None:
|
||||
return True
|
||||
get_output_embeddings = getattr(module, "get_output_embeddings", None)
|
||||
if callable(get_output_embeddings):
|
||||
try:
|
||||
return get_output_embeddings() is not None
|
||||
except Exception:
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
def _grpo_causal_head(model):
|
||||
# The module that owns lm_head / output embeddings (whose forward emits `.logits`).
|
||||
get_base_model = getattr(model, "get_base_model", None)
|
||||
if callable(get_base_model):
|
||||
base_model = get_base_model()
|
||||
if base_model is not None:
|
||||
return base_model
|
||||
return model
|
||||
|
||||
|
||||
def _grpo_hidden_states_wrap_target(model):
|
||||
if model is None:
|
||||
return None
|
||||
|
|
@ -588,10 +613,14 @@ def _grpo_hidden_states_wrap_target(model):
|
|||
base_model = get_base_model()
|
||||
if base_model is not None and base_model is not model:
|
||||
return base_model
|
||||
for attr in ("base_model", "model"):
|
||||
child = getattr(model, attr, None)
|
||||
if child is not None and child is not model and hasattr(child, "forward"):
|
||||
return child
|
||||
# Only descend into a child when `model` does not own lm_head itself. GRPO consumes the
|
||||
# `.logits` of the module that applies lm_head; wrapping the inner trunk (no lm_head) would
|
||||
# let the outer forward re-apply lm_head and leak logits into the chunked log-softmax (#708).
|
||||
if not _grpo_owns_lm_head(model):
|
||||
for attr in ("base_model", "model"):
|
||||
child = getattr(model, attr, None)
|
||||
if child is not None and child is not model and hasattr(child, "forward"):
|
||||
return child
|
||||
return model
|
||||
|
||||
|
||||
|
|
@ -686,12 +715,53 @@ def _replace_outputs_logits(outputs, hidden_states):
|
|||
)
|
||||
|
||||
|
||||
def _install_grpo_lm_head_passthrough(model):
|
||||
# Preferred hidden-states path for a plain *ForCausalLM (e.g. full finetuning, where the model
|
||||
# is not PEFT-wrapped, keeps the stock HF forward, and so has no RETURN_HIDDEN_STATES branch or
|
||||
# support marker). Short-circuit lm_head to return its input (the hidden states) when
|
||||
# UNSLOTH_RETURN_HIDDEN_STATES=1; the forward then yields `.logits == hidden`, which the GRPO
|
||||
# log-prob path projects in chunks itself, and the full vocab projection is skipped. The lm_head
|
||||
# weight is untouched, and the accelerate-managed top-level forward is not wrapped, so there is
|
||||
# no bound-self collision. No-op when the flag is 0.
|
||||
head = _grpo_causal_head(model)
|
||||
lm_head = getattr(head, "lm_head", None)
|
||||
if lm_head is None:
|
||||
get_output_embeddings = getattr(head, "get_output_embeddings", None)
|
||||
if callable(get_output_embeddings):
|
||||
try:
|
||||
lm_head = get_output_embeddings()
|
||||
except Exception:
|
||||
lm_head = None
|
||||
if lm_head is None or getattr(lm_head, "_unsloth_grpo_passthrough", False):
|
||||
return False
|
||||
|
||||
original_lm_head_forward = lm_head.forward
|
||||
|
||||
def passthrough_forward(*args, **kwargs):
|
||||
if os.environ.get("UNSLOTH_RETURN_HIDDEN_STATES", "0") == "1":
|
||||
return args[0] if args else next(iter(kwargs.values()))
|
||||
return original_lm_head_forward(*args, **kwargs)
|
||||
|
||||
lm_head.forward = passthrough_forward
|
||||
lm_head._unsloth_grpo_passthrough = True
|
||||
setattr(model, _UNSLOTH_RETURN_HIDDEN_STATES_SUPPORT_MARKER, True)
|
||||
setattr(head, _UNSLOTH_RETURN_HIDDEN_STATES_SUPPORT_MARKER, True)
|
||||
return True
|
||||
|
||||
|
||||
def _install_grpo_hidden_states_forward_wrapper(model):
|
||||
if model is None or getattr(model, _UNSLOTH_GRPO_HIDDEN_STATES_WRAPPED_ATTR, False):
|
||||
return False
|
||||
if _model_supports_unsloth_return_hidden_states(model):
|
||||
return False
|
||||
|
||||
# Preferred: short-circuit lm_head (robust for a plain full-FT CausalLM, skips the vocab
|
||||
# projection, and avoids wrapping the accelerate-managed top-level forward). Fall back to the
|
||||
# forward wrapper only when no lm_head can be found.
|
||||
if _install_grpo_lm_head_passthrough(model):
|
||||
setattr(model, _UNSLOTH_GRPO_HIDDEN_STATES_WRAPPED_ATTR, True)
|
||||
return True
|
||||
|
||||
target_model = _grpo_hidden_states_wrap_target(model)
|
||||
if getattr(target_model, _UNSLOTH_GRPO_HIDDEN_STATES_WRAPPED_ATTR, False):
|
||||
setattr(model, _UNSLOTH_GRPO_HIDDEN_STATES_WRAPPED_ATTR, True)
|
||||
|
|
@ -702,6 +772,10 @@ def _install_grpo_hidden_states_forward_wrapper(model):
|
|||
model_name = type(target_model).__name__
|
||||
|
||||
def wrapped_forward(*args, **kwargs):
|
||||
# Tolerate being invoked as a bound method: accelerate / nn.Module __call__ can inject
|
||||
# `self` as the first positional arg once the wrapper lives on the outer CausalLM.
|
||||
if len(args) > 0 and args[0] is target_model:
|
||||
args = args[1:]
|
||||
if os.environ.get("UNSLOTH_RETURN_HIDDEN_STATES", "0") != "1":
|
||||
return original_forward(*args, **kwargs)
|
||||
|
||||
|
|
@ -1027,8 +1101,18 @@ def _patch_trl_rl_trainers_impl(trainer_file = "grpo_trainer"):
|
|||
"use_fp16 = getattr(args, 'fp16', False)\n"
|
||||
"if type(use_fp16) is not bool: use_fp16 = False\n"
|
||||
"force_float32 = False\n"
|
||||
# device-aware bf16 check (CUDA/XPU/HIP), so V100/T4 never pick bf16
|
||||
# but AMD/Intel are unaffected; fall back on older unsloth_zoo.
|
||||
"try:\n"
|
||||
" from unsloth_zoo.device_type import device_is_bf16_supported as _bf16_supported\n"
|
||||
"except Exception:\n"
|
||||
" _bf16_supported = torch.cuda.is_bf16_supported\n"
|
||||
# FORCE_FLOAT32 models (Gemma3, gpt_oss, ...) cannot use float16. On a GPU without
|
||||
# bf16 (V100/T4) keep them in float32 so they never autocast to fp16. On a bf16 GPU,
|
||||
# full finetuning can still use bf16 autocast (master weights stay float32), which is
|
||||
# faster and uses less memory; LoRA/QLoRA keep float32 when forced.
|
||||
"full_finetuning = os.environ.get('UNSLOTH_ENABLE_FULL_FINETUNING', '0') == '1'\n"
|
||||
"if not full_finetuning and (os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1'):\n"
|
||||
"if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1' and not (full_finetuning and _bf16_supported()):\n"
|
||||
" print('Unsloth: Switching to float32 training since model cannot work with float16')\n"
|
||||
" force_float32 = True\n"
|
||||
"mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')\n"
|
||||
|
|
@ -1037,8 +1121,9 @@ def _patch_trl_rl_trainers_impl(trainer_file = "grpo_trainer"):
|
|||
"from unsloth_zoo.utils import _get_dtype\n"
|
||||
"dtype = _get_dtype(dtype)\n"
|
||||
"float16 = dtype == torch.float16\n"
|
||||
"bfloat16 = dtype == torch.bfloat16\n"
|
||||
"if not force_float32 and (float16 and use_bf16): raise TypeError('Unsloth: Model is in float16 precision but you want to use bfloat16 precision. Set fp16 to `True` and bf16 to `False`')\n"
|
||||
"if not force_float32 and (not float16 and use_fp16): raise TypeError('Unsloth: Model is in bfloat16 precision but you want to use float16 precision. Set fp16 to `False` and bf16 to `True`')\n"
|
||||
"if not force_float32 and (bfloat16 and use_fp16): raise TypeError('Unsloth: Model is in bfloat16 precision but you want to use float16 precision. Set fp16 to `False` and bf16 to `True`')\n"
|
||||
"if force_float32:\n"
|
||||
" # Forced float32 training\n"
|
||||
" args.fp16 = False\n"
|
||||
|
|
@ -1047,11 +1132,12 @@ def _patch_trl_rl_trainers_impl(trainer_file = "grpo_trainer"):
|
|||
" if hasattr(args, 'mixed_precision'): args.mixed_precision = 'no'\n"
|
||||
" # args.mixed_precision is a new argument which needs to be set now\n"
|
||||
"elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':\n"
|
||||
" # Mixed precision training\n"
|
||||
" args.fp16 = float16\n"
|
||||
" args.bf16 = not float16\n"
|
||||
" os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'\n"
|
||||
" if hasattr(args, 'mixed_precision'): args.mixed_precision = 'fp16' if float16 else 'bf16'\n"
|
||||
" # Mixed precision training. bf16 only if the GPU supports it; V100/T4 use fp16.\n"
|
||||
" use_bf16_amp = (not float16) and _bf16_supported()\n"
|
||||
" args.fp16 = not use_bf16_amp\n"
|
||||
" args.bf16 = use_bf16_amp\n"
|
||||
" os.environ['ACCELERATE_MIXED_PRECISION'] = 'bf16' if use_bf16_amp else 'fp16'\n"
|
||||
" if hasattr(args, 'mixed_precision'): args.mixed_precision = 'bf16' if use_bf16_amp else 'fp16'\n"
|
||||
" # args.mixed_precision is a new argument which needs to be set now\n"
|
||||
"elif mixed_precision_dtype == 'bfloat16':\n"
|
||||
" # Both False since bfloat16 full finetuning doesn't do any autocasting.\n"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue