diff --git a/tests/version_compat/test_trl_grpo_pinned_symbols.py b/tests/version_compat/test_trl_grpo_pinned_symbols.py index 8f1435ba5f..4c7dcc4234 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): - """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.""" + """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.""" candidates = [ "trl/trainer/kto_trainer.py", "trl/experimental/kto/kto_trainer.py", @@ -566,15 +566,11 @@ 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}: KTO log-prob computation not found in any of {candidates}; " - f"unsloth/models/rl_replacements.py KTO rewrite silently skipped" + f"{tag}: KTOTrainer.get_batch_logps not found in any of {candidates}; " + f"unsloth/models/rl_replacements.py:1675 rewrite silently skipped" ) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index c1d92a31c4..31d54675c9 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -1962,63 +1962,6 @@ 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):