From a1e2333d7c4c9532691f9275b8593d5285f060d1 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 31 May 2026 12:58:16 +0000 Subject: [PATCH 1/2] Fix GRPO full finetuning returning logits instead of hidden states In full finetuning the model reaches GRPO as a plain *ForCausalLM with the stock forward, so it has no UNSLOTH_RETURN_HIDDEN_STATES branch and no support marker. The fallback then mis-targets the inner trunk (no lm_head) via _grpo_hidden_states_wrap_target, so the outer forward re-applies lm_head and the chunked log-softmax receives logits (vocab) instead of hidden states (hidden), crashing in the lm_head matmul. Prefer an lm_head passthrough: when UNSLOTH_RETURN_HIDDEN_STATES=1, short circuit lm_head to return its input so the forward yields logits == hidden and the vocab projection is skipped (memory efficient). Also harden the forward-wrapper fallback for models without a discoverable lm_head: do not descend past the lm_head owner, and tolerate an injected leading module arg from accelerate. No-op for LoRA/QLoRA and when the flag is unset. --- unsloth/models/rl.py | 78 +++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 74 insertions(+), 4 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 359716fda4..8374920e75 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,49 @@ 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 +768,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) From ca826bc3d8f712eaf24a0fe698285746ad167e62 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sun, 31 May 2026 12:59:12 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/rl.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 8374920e75..82ef7b48d2 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -728,16 +728,20 @@ def _install_grpo_lm_head_passthrough(model): 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 + 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)