test(version-compat): keep GRPO fake-run logits finite on CPU (#7247)

* test(version-compat): keep GRPO fake-run logits finite on CPU

The GRPO fake-run test samples completions from a tiny untrained model on
CPU. Such a model can emit non-finite logits, so torch.multinomial inside
generate() intermittently raises "probability tensor contains either inf,
nan or element < 0" -- a nondeterministic sampling failure, not a regression
(the Trainer already fixes the seed, but CPU reduction order is not
bit-reproducible). Add a forward hook that sanitizes the LM head logits to a
finite bounded range before sampling, so the fake run reliably exercises the
whole train loop; the test checks the loop runs, not the numerics.

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

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

* test(version-compat): drop redundant nan_to_num bounds (clamp handles them)

* test(version-compat): scope GRPO finite-logits guard to the GRPO test

Only test_grpo_trains_on_cpu autoregressively samples completions, so it is
the only canary that can hit the non-finite-logits torch.multinomial crash.
Move the _guard_finite_logits hook out of the shared _load_plain() and into
test_grpo_trains_on_cpu so the SFT and DPO canaries keep asserting against the
model's true, unclamped logits.

---------

Co-authored-by: Daniel Han <unslothai@gmail.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Daniel Han 2026-07-19 04:54:17 -07:00 committed by GitHub
commit ecd97a935a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -158,6 +158,37 @@ except Exception:
_MODEL = "hf-internal-testing/tiny-random-LlamaForCausalLM" _MODEL = "hf-internal-testing/tiny-random-LlamaForCausalLM"
def _guard_finite_logits(model):
"""Keep the LM head logits finite so GRPO sampling can't crash.
``test_grpo_trains_on_cpu`` samples completions from a tiny, *untrained*
random model on CPU. Driven autoregressively -- and nudged by the fake
reward's optimizer step between the two train steps -- such a model can emit
non-finite logits, so ``torch.multinomial`` inside ``generate()``
intermittently raises "probability tensor contains either `inf`, `nan` or
element < 0". That is a well-known nondeterministic sampling failure, not an
Unsloth/TRL regression: the Trainer already fixes the seed, but CPU reduction
order is not bit-reproducible, so the blow-up still surfaces every so often.
Sanitize the logits to a finite, bounded range (out of place, so autograd
stays valid) before they reach the sampler. This test asserts the train loop
runs end to end, not the (deliberately meaningless) numerics, so bounding the
logits changes nothing it checks while making the run reliable.
"""
def _finite_logits_hook(_module, _inputs, output):
logits = getattr(output, "logits", None)
if logits is None:
return output
# nan_to_num maps nan -> 0 and the infinities to large finite values;
# clamp then bounds everything to [-30, 30].
output.logits = torch.nan_to_num(logits).clamp(-30.0, 30.0)
return output
model.register_forward_hook(_finite_logits_hook)
return model
def _load_plain(): def _load_plain():
"""Tiny plain HF model + tokenizer on CPU. Skips (not fails) if the model """Tiny plain HF model + tokenizer on CPU. Skips (not fails) if the model
cannot be fetched -- that is a network/hub issue, not an unsloth regression.""" cannot be fetched -- that is a network/hub issue, not an unsloth regression."""
@ -233,6 +264,11 @@ def test_grpo_trains_on_cpu(tmp_path):
assert GRPOTrainer.__name__ == "UnslothGRPOTrainer", "GRPO patch did not apply" assert GRPOTrainer.__name__ == "UnslothGRPOTrainer", "GRPO patch did not apply"
model, tok = _load_plain() model, tok = _load_plain()
# GRPO is the only canary that autoregressively samples completions, so it is
# the only one that can hit the non-finite-logits multinomial crash. Install
# the guard here (not in _load_plain) so the SFT/DPO canaries keep asserting
# against the model's true, unclamped outputs.
_guard_finite_logits(model)
ds = Dataset.from_list([{"prompt": "hi there"}] * 4) ds = Dataset.from_list([{"prompt": "hi there"}] * 4)
cfg = GRPOConfig( cfg = GRPOConfig(
output_dir = str(tmp_path / "ci_grpo"), output_dir = str(tmp_path / "ci_grpo"),