From 157cecb25c3c7277a6da33c81a409373adb2b4b1 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 4 Jun 2026 00:07:58 -0700 Subject: [PATCH] Port KTO logps truncation guard to TRL 1.x _compute_logps refactor (#5996) * Port KTO logps truncation guard to TRL 1.x _compute_logps refactor * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --- .../test_trl_grpo_pinned_symbols.py | 20 ++++--- unsloth/models/rl_replacements.py | 57 +++++++++++++++++++ 2 files changed, 69 insertions(+), 8 deletions(-) diff --git a/tests/version_compat/test_trl_grpo_pinned_symbols.py b/tests/version_compat/test_trl_grpo_pinned_symbols.py index 4c7dcc4234..8f1435ba5f 100644 --- a/tests/version_compat/test_trl_grpo_pinned_symbols.py +++ b/tests/version_compat/test_trl_grpo_pinned_symbols.py @@ -551,12 +551,12 @@ def test_trl_grpo_source_inference_mode_unwrap(tag: str): @pytest.mark.parametrize("tag", TRL_TAGS) def test_trl_kto_get_batch_logps_signature(tag: str): - """TRL 0.27+ moved KTOTrainer to trl.experimental.kto and the - canonical kto_trainer.py shrank to a thin re-export wrapper. The - real `get_batch_logps` lives at trl/experimental/kto/kto_trainer.py. - Unsloth's MRO walk in models/rl.py:592-708 already follows - trl.experimental.* parents, so either path is fine — we just - require the symbol to exist SOMEWHERE.""" + """KTO log-prob computation must stay patchable. Through TRL 1.x the + target was KTOTrainer.get_batch_logps; TRL 1.x dropped it and moved the + math into _compute_logps / compute_ref_log_probs calling + selective_log_softmax. unsloth/models/rl_replacements.py patches BOTH + shapes (kto_trainer_get_batch_logps + kto_trainer_align_completion_logps), + so we require EITHER form to exist wherever KTOTrainer lives.""" candidates = [ "trl/trainer/kto_trainer.py", "trl/experimental/kto/kto_trainer.py", @@ -566,11 +566,15 @@ def test_trl_kto_get_batch_logps_signature(tag: str): src = fetch_text("huggingface/trl", tag, path) if src is None: continue + # Legacy: explicit get_batch_logps method. if has_def(src, "get_batch_logps", "func"): return + # TRL 1.x: refactored into _compute_logps + selective_log_softmax. + if has_def(src, "_compute_logps", "func") and "selective_log_softmax" in src: + return pytest.fail( - f"{tag}: KTOTrainer.get_batch_logps not found in any of {candidates}; " - f"unsloth/models/rl_replacements.py:1675 rewrite silently skipped" + f"{tag}: KTO log-prob computation not found in any of {candidates}; " + f"unsloth/models/rl_replacements.py KTO rewrite silently skipped" ) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 31d54675c9..c1d92a31c4 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -1962,6 +1962,63 @@ def kto_trainer_get_batch_logps(function_name, function): RL_FUNCTIONS["kto_trainer"].append(kto_trainer_get_batch_logps) +# TRL 1.x dropped KTOTrainer.get_batch_logps and moved the log-prob math into +# _compute_logps / compute_ref_log_probs / _compute_kl_logps, which call +# selective_log_softmax on completion-only tokens. Same truncation hazard as +# above, so clamp logits/ids/mask to the shorter seq length (no-op when equal). +_KTO_COMPLETION_RE = re.compile( + r"(?P[ \t]*)shift_logits = completion_logits\[:, :-1, :\]\.contiguous\(\)\n" + r"(?P=ws)per_token_logps = selective_log_softmax\(\s*shift_logits,\s*" + r"(?P\w+)\[\"completion_input_ids\"\]\[:, 1:\]\.contiguous\(\)\s*\)\n" + r"(?P=ws)per_token_logps\[(?P=var)\[\"completion_mask\"\]\[:, 1:\] == 0\] = 0\.0" +) +_KTO_KL_RE = re.compile( + r"(?P[ \t]*)shift_KL_logits = KL_logits\[:, :-1, :\]\.contiguous\(\)\n" + r"(?P=ws)KL_per_token_logps = selective_log_softmax\(\s*shift_KL_logits,\s*" + r"(?P\w+)\[\"KL_completion_input_ids\"\]\[:, 1:\]\.contiguous\(\)\s*\)\n" + r"(?P=ws)KL_per_token_logps\[(?P=var)\[\"KL_completion_mask\"\]\[:, 1:\] == 0\] = 0\.0" +) + + +def _kto_completion_repl(m): + ws, var = m.group("ws"), m.group("var") + return ( + f"{ws}shift_logits = completion_logits[:, :-1, :].contiguous()\n" + f"{ws}# Unsloth: clamp logits/ids/mask to shorter seq len (model may truncate input_ids)\n" + f'{ws}_uns_ids = {var}["completion_input_ids"][:, 1:].contiguous()\n' + f"{ws}_uns_n = min(shift_logits.shape[1], _uns_ids.shape[1])\n" + f"{ws}per_token_logps = selective_log_softmax(shift_logits[:, :_uns_n], _uns_ids[:, :_uns_n])\n" + f'{ws}per_token_logps[{var}["completion_mask"][:, 1:][:, :_uns_n] == 0] = 0.0' + ) + + +def _kto_kl_repl(m): + ws, var = m.group("ws"), m.group("var") + return ( + f"{ws}shift_KL_logits = KL_logits[:, :-1, :].contiguous()\n" + f"{ws}# Unsloth: clamp logits/ids/mask to shorter seq len (model may truncate input_ids)\n" + f'{ws}_uns_kl_ids = {var}["KL_completion_input_ids"][:, 1:].contiguous()\n' + f"{ws}_uns_kl_n = min(shift_KL_logits.shape[1], _uns_kl_ids.shape[1])\n" + f"{ws}KL_per_token_logps = selective_log_softmax(shift_KL_logits[:, :_uns_kl_n], _uns_kl_ids[:, :_uns_kl_n])\n" + f'{ws}KL_per_token_logps[{var}["KL_completion_mask"][:, 1:][:, :_uns_kl_n] == 0] = 0.0' + ) + + +def kto_trainer_align_completion_logps(function_name, function): + if function_name not in ( + "_compute_logps", + "compute_ref_log_probs", + "_compute_kl_logps", + ): + return function + function = _KTO_COMPLETION_RE.sub(_kto_completion_repl, function) + function = _KTO_KL_RE.sub(_kto_kl_repl, function) + return function + + +RL_FUNCTIONS["kto_trainer"].append(kto_trainer_align_completion_logps) + + # https://github.com/huggingface/trl/blob/main/trl/trainer/grpo_trainer.py#L356 # TRL warns if batch size is not a multiple of num_generations -> fix this. def grpo_trainer_fix_batch_size(RLTrainer_source, RLConfig_source):