From a1e2333d7c4c9532691f9275b8593d5285f060d1 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 31 May 2026 12:58:16 +0000 Subject: [PATCH 1/4] 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/4] [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) From 11769cf006ef349fdd451cf777b213614e259792 Mon Sep 17 00:00:00 2001 From: Datta Nimmaturi Date: Tue, 2 Jun 2026 10:59:09 +0000 Subject: [PATCH 3/4] Use context instead of manual setting for UNSLOTH_RETURN_HIDDEN_STATES --- unsloth/models/rl.py | 66 +++++++++++++++++++++++++++++-- unsloth/models/rl_replacements.py | 31 +++++++++++---- 2 files changed, 87 insertions(+), 10 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 82ef7b48d2..0768915a73 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -597,10 +597,14 @@ def _grpo_owns_lm_head(module): def _grpo_causal_head(model): # The module that owns lm_head / output embeddings (whose forward emits `.logits`). + if model is None: + return None + if _grpo_owns_lm_head(model): + return model 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: + if base_model is not None and _grpo_owns_lm_head(base_model): return base_model return model @@ -608,6 +612,8 @@ def _grpo_causal_head(model): def _grpo_hidden_states_wrap_target(model): if model is None: return None + if _grpo_owns_lm_head(model): + return model get_base_model = getattr(model, "get_base_model", None) if callable(get_base_model): base_model = get_base_model() @@ -624,6 +630,53 @@ def _grpo_hidden_states_wrap_target(model): return model +def _grpo_config_value(config, name, default = None): + if config is None: + return default + value = getattr(config, name, default) + if value is not default: + return value + text_config = getattr(config, "text_config", None) + if text_config is not None: + value = getattr(text_config, name, default) + if value is not default: + return value + get_text_config = getattr(config, "get_text_config", None) + if callable(get_text_config): + try: + text_config = get_text_config() + return getattr(text_config, name, default) + except Exception: + pass + return default + + +def _grpo_has_post_lm_head_transform(module): + config = getattr(module, "config", None) + if config is None: + return False + + final_logit_softcapping = _grpo_config_value( + config, "final_logit_softcapping", None + ) + if final_logit_softcapping not in (None, 0, 0.0): + return True + + logit_scale = _grpo_config_value(config, "logit_scale", None) + if logit_scale not in (None, 0, 0.0, 1, 1.0): + return True + + logits_scaling = _grpo_config_value(config, "logits_scaling", None) + if logits_scaling not in (None, 0, 0.0, 1, 1.0): + return True + + lm_head_multiplier = _grpo_config_value(config, "lm_head_multiplier", None) + if lm_head_multiplier not in (None, 0, 0.0, 1, 1.0): + return True + + return False + + def _model_supports_unsloth_return_hidden_states(model): target_model = _grpo_hidden_states_wrap_target(model) for candidate in (model, target_model): @@ -724,6 +777,8 @@ def _install_grpo_lm_head_passthrough(model): # 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) + if _grpo_has_post_lm_head_transform(head): + return False lm_head = getattr(head, "lm_head", None) if lm_head is None: get_output_embeddings = getattr(head, "get_output_embeddings", None) @@ -738,9 +793,14 @@ def _install_grpo_lm_head_passthrough(model): original_lm_head_forward = lm_head.forward def passthrough_forward(*args, **kwargs): + forward_args = args[1:] if len(args) > 0 and args[0] is lm_head else args 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) + if len(forward_args) > 0: + return forward_args[0] + if len(kwargs) > 0: + return next(iter(kwargs.values())) + raise TypeError("forward() missing 1 required positional argument: 'input'") + return original_lm_head_forward(*forward_args, **kwargs) lm_head.forward = passthrough_forward lm_head._unsloth_grpo_passthrough = True diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 31d54675c9..a0e088af49 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -25,6 +25,7 @@ import re import torch import inspect import linecache +from contextlib import contextmanager from collections import defaultdict from unsloth_zoo.rl_replacements import ( RL_REPLACEMENTS, @@ -74,6 +75,19 @@ RL_CONFIG_CHANGES = defaultdict(list) RL_METRICS_CHANGES = defaultdict(list) RL_ADDITIONAL_FUNCTIONS = defaultdict(list) + +@contextmanager +def _temporary_unsloth_return_hidden_states(): + old_value = os.environ.get("UNSLOTH_RETURN_HIDDEN_STATES") + os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" + try: + yield + finally: + if old_value is None: + os.environ.pop("UNSLOTH_RETURN_HIDDEN_STATES", None) + else: + os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = old_value + _DPO_VISION_KEYS = ( "pixel_position_ids", "image_position_ids", @@ -1100,8 +1114,9 @@ def grpo_trainer__get_per_token_logps(function_name, function): if os.environ.get("UNSLOTH_FORCE_FLOAT32", "0") == "1": self._autocast_dtype = torch.float16 - os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" - with torch.amp.autocast(device_type = DEVICE_TYPE, dtype = self._autocast_dtype): + with _temporary_unsloth_return_hidden_states(), torch.amp.autocast( + device_type = DEVICE_TYPE, dtype = self._autocast_dtype + ): # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded logits = model( input_ids = input_ids, @@ -1379,9 +1394,9 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): token_type_ids_chunks, mm_token_type_ids_chunks, ) - os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" - - with _get_inference_mode_context_manager(model): + with _temporary_unsloth_return_hidden_states(), _get_inference_mode_context_manager( + model + ): for ( input_ids_chunk, attention_mask_chunk, @@ -1478,8 +1493,6 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): logprobs = torch.cat(all_logprobs_list, dim = 0) entropies = None - os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "0" - return logprobs.detach(), entropies # logps, entropies # input_ids = input_ids[:, -logits_to_keep:] # For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves. @@ -1536,6 +1549,10 @@ grpo_update_SamplingParams = RL_REPLACEMENTS["grpo_update_SamplingParams"] RL_PRE_ITEMS["grpo_trainer"].append( inspect.getsource(_unsloth_get_final_logit_softcapping) ) +RL_PRE_ITEMS["grpo_trainer"].append("from contextlib import contextmanager") +RL_PRE_ITEMS["grpo_trainer"].append( + inspect.getsource(_temporary_unsloth_return_hidden_states) +) RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(_unsloth_get_mm_token_id)) RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(_unsloth_fix_mm_token_type_ids)) RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(_unsloth_clear_stateful_mrope)) From 82e1468dcb0f9b5567c6b6c88d88d0d31b8419ad Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Tue, 2 Jun 2026 10:59:56 +0000 Subject: [PATCH 4/4] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/models/rl_replacements.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index a0e088af49..774e0dca5b 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -88,6 +88,7 @@ def _temporary_unsloth_return_hidden_states(): else: os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = old_value + _DPO_VISION_KEYS = ( "pixel_position_ids", "image_position_ids", @@ -1114,8 +1115,9 @@ def grpo_trainer__get_per_token_logps(function_name, function): if os.environ.get("UNSLOTH_FORCE_FLOAT32", "0") == "1": self._autocast_dtype = torch.float16 - with _temporary_unsloth_return_hidden_states(), torch.amp.autocast( - device_type = DEVICE_TYPE, dtype = self._autocast_dtype + with ( + _temporary_unsloth_return_hidden_states(), + torch.amp.autocast(device_type = DEVICE_TYPE, dtype = self._autocast_dtype), ): # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded logits = model( @@ -1394,8 +1396,9 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function): token_type_ids_chunks, mm_token_type_ids_chunks, ) - with _temporary_unsloth_return_hidden_states(), _get_inference_mode_context_manager( - model + with ( + _temporary_unsloth_return_hidden_states(), + _get_inference_mode_context_manager(model), ): for ( input_ids_chunk,