diff --git a/tests/python/test_v100_fullft_precision.py b/tests/python/test_v100_fullft_precision.py new file mode 100644 index 0000000000..f15b574f23 --- /dev/null +++ b/tests/python/test_v100_fullft_precision.py @@ -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" diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 359716fda4..7acac36b50 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -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"