Compare commits
5 commits
main
...
grpo-fullf
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1b484332ba |
||
|
|
82e1468dcb | ||
|
|
11769cf006 | ||
|
|
ca826bc3d8 | ||
|
|
a1e2333d7c |
2 changed files with 165 additions and 11 deletions
|
|
@ -580,21 +580,103 @@ _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`).
|
||||
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 and _grpo_owns_lm_head(base_model):
|
||||
return base_model
|
||||
return 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()
|
||||
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
|
||||
|
||||
|
||||
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):
|
||||
|
|
@ -686,12 +768,60 @@ 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)
|
||||
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)
|
||||
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):
|
||||
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":
|
||||
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
|
||||
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 +832,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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,20 @@ 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 +1115,10 @@ 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 +1396,10 @@ 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 +1496,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 +1552,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))
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue