fast_generate: unify legacy/new logits kwarg + fix Mistral merge site (#5543)

* fast_generate: unify legacy/new logits kwarg + fix Mistral merge site

Two related issues caught by review on PR #5538:

1. unsloth_fast_generate (models/llama.py)

   The previous patch promoted num_logits_to_keep -> logits_to_keep
   unconditionally whenever the caller supplied num_logits_to_keep,
   and only popped num_logits_to_keep (not logits_to_keep). On
   transformers older than 4.50 (legacy spelling is the only one the
   model forward accepts), the promotion broke things; symmetrically,
   a caller supplying logits_to_keep on those older transformers also
   went unchecked.

   Switch to the unified normalize-then-inspect pattern from the
   review:

     _provided_num    = kwargs.pop("num_logits_to_keep", None)
     _provided_logits = kwargs.pop("logits_to_keep",     None)
     _provided = _provided_logits if _provided_logits is not None else _provided_num
     _fwd_params = inspect.signature(self.forward).parameters
     if "logits_to_keep" in _fwd_params:
         kwargs["logits_to_keep"] = _provided if _provided is not None else 1
     elif "num_logits_to_keep" in _fwd_params:
         kwargs["num_logits_to_keep"] = _provided if _provided is not None else 1

   Inspect the runtime forward signature first, then choose the
   spelling it actually accepts, then route either user-supplied value
   under that spelling. Backward-compatible in both directions.

2. MistralForCausalLM_fast_forward (models/mistral.py)

   The max(num_logits_to_keep, logits_to_keep) merge was inside the
   `if UNSLOTH_RETURN_HIDDEN_STATES:` block, so it only fired on the
   GRPO hidden-states path. On the normal generation path the elif at
   line 316 only checked num_logits_to_keep, so a caller (including
   unsloth_fast_generate itself) passing logits_to_keep=1 ended up
   computing full prompt logits instead of slicing to the last token.
   For long prompts that reintroduces the large prefill logits
   allocation the default keep=1 was avoiding.

   Move the max() merge above the env-var branching so the normal
   generation path slices correctly too. Llama already did this
   merge at the top (unsloth/models/llama.py:1501); Mistral now
   matches.

No behaviour change on the default GRPO / SFT paths. Targets only the
edge cases the review flagged.

* fast_generate: preserve caller logits kwarg when signature inspect fails

If `inspect.signature(self.forward)` raises TypeError/ValueError (opaque
C-extension or compiled wrappers), the previous fix set `_fwd_params = {}`
which silently dropped the caller-supplied `logits_to_keep` /
`num_logits_to_keep`. Fall back to the spelling the caller used (default
`logits_to_keep=1` when neither was supplied) so generation still honors
the requested logits slice.

* fast_forward: do not max() int against tensor logits_to_keep

HF accepts logits_to_keep as a 1-D LongTensor of positions for
selective decode. The merge in mistral.py (added by this PR) and
the pre-existing one in llama.py both run max(int, Tensor), which
casts the comparison to a bool and raises on multi-element tensors.
Branch on type and skip the merge when either argument is a tensor;
downstream int-slice path is unchanged, so tensor callers fall
through with num_logits_to_keep == 0, matching pre-merge behavior.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* fast_generate/forward: shorten kwarg-merge comments

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Daniel Han 2026-05-18 04:19:48 -07:00 committed by GitHub
commit 3ebe17fe41
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 34 additions and 20 deletions

View file

@ -1498,7 +1498,13 @@ def CausalLM_fast_forward(fast_forward_inference):
logit_softcapping = getattr(self.config, "final_logit_softcapping", 0)
logit_scaling = getattr(self.config, "logit_scale", 0)
dtype = lm_head.dtype
num_logits_to_keep = max(num_logits_to_keep, logits_to_keep)
# Skip int max() if either is a tensor (HF selective-decode form).
if isinstance(num_logits_to_keep, torch.Tensor) or isinstance(
logits_to_keep, torch.Tensor
):
num_logits_to_keep = 0
else:
num_logits_to_keep = max(num_logits_to_keep, logits_to_keep)
# Move items to same device as lm_head
hidden_states = hidden_states.to(lm_head_device)
@ -2109,24 +2115,23 @@ def unsloth_fast_generate(
# For newer HF
kwargs["cache_implementation"] = "dynamic"
# transformers 4.50 renamed num_logits_to_keep -> logits_to_keep
# (with @deprecate_kwarg through 4.51.x, removed in 4.52+). Pick the
# spelling the actual runtime forward accepts so generation
# _validate_model_kwargs does not reject the legacy name.
num_logits_to_keep = kwargs.pop("num_logits_to_keep", None)
logits_to_keep = kwargs.get("logits_to_keep", None)
if num_logits_to_keep is not None and logits_to_keep is None:
kwargs["logits_to_keep"] = num_logits_to_keep
logits_to_keep = num_logits_to_keep
if num_logits_to_keep is None and logits_to_keep is None:
try:
_fwd_params = inspect.signature(self.forward).parameters
except (TypeError, ValueError):
_fwd_params = {}
if "logits_to_keep" in _fwd_params:
kwargs["logits_to_keep"] = 1
elif "num_logits_to_keep" in _fwd_params:
kwargs["num_logits_to_keep"] = 1
# transformers 4.50 renamed num_logits_to_keep -> logits_to_keep.
# Pop both, re-emit under the spelling forward() accepts.
_provided_num = kwargs.pop("num_logits_to_keep", None)
_provided_logits = kwargs.pop("logits_to_keep", None)
_provided = _provided_logits if _provided_logits is not None else _provided_num
try:
_fwd_params = inspect.signature(self.forward).parameters
_has_new = "logits_to_keep" in _fwd_params
_has_old = "num_logits_to_keep" in _fwd_params
except (TypeError, ValueError):
# Opaque forward: keep the caller's spelling, default to new.
_has_old = _provided_num is not None and _provided_logits is None
_has_new = not _has_old
if _has_new:
kwargs["logits_to_keep"] = _provided if _provided is not None else 1
elif _has_old:
kwargs["num_logits_to_keep"] = _provided if _provided is not None else 1
# Remove token_type_ids
kwargs.pop("token_type_ids", None)

View file

@ -297,9 +297,18 @@ def MistralForCausalLM_fast_forward(
if labels is not None:
labels = labels.to(lm_head_device)
# Merge legacy / new spellings before branching so the decode-time
# last-token slice fires on the normal path too. Skip int max() if
# either is a tensor (HF selective-decode form).
if isinstance(num_logits_to_keep, torch.Tensor) or isinstance(
logits_to_keep, torch.Tensor
):
num_logits_to_keep = 0
else:
num_logits_to_keep = max(num_logits_to_keep, logits_to_keep)
# If we are in GRPO mode, return raw hidden states
if os.environ.get("UNSLOTH_RETURN_HIDDEN_STATES", "0") == "1":
num_logits_to_keep = max(num_logits_to_keep, logits_to_keep)
if num_logits_to_keep != 0:
hidden_states = hidden_states[:, -num_logits_to_keep:, :]
return CausalLMOutputWithPast(