From 114c91ed030e6d6f00b9fd3f5d7fdd753cf0bace Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 22 Jun 2025 05:51:07 -0700 Subject: [PATCH] Update rl_replacements.py --- unsloth/models/rl_replacements.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 5d4fc57b77..18f7720562 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -263,8 +263,6 @@ def grpo_trainer__get_per_token_logps(function_name, function): # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded hidden_states = model(input_ids=input_ids, attention_mask=attention_mask, logits_to_keep=logits_to_keep + 1).logits #logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred - print("##############, input_ids", input_ids.shape) - print("##############, hidden_states", hidden_states.shape) return hidden_states # input_ids = input_ids[:, -logits_to_keep:] # For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves. @@ -339,9 +337,7 @@ def grpo_trainer_compute_loss(function_name, function): else: old_hidden_states = None - print("$$$$$$$$$$$ input_ids", input_ids.shape) input_ids = input_ids[:, -logits_to_keep:] - print("$$$$$$$$$$$ input_ids", input_ids.shape) if per_token_logps is not None: if ref_per_token_logps is not None: