From 3fa2f1436a10daf5c723c9ab62ccb5efb5324a44 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 22 Jun 2025 02:37:52 -0700 Subject: [PATCH] Update rl_replacements.py --- unsloth/models/rl_replacements.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index df95f73fc5..e343a2ec70 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -320,12 +320,16 @@ def grpo_trainer_compute_loss(function_name, function): # _prepare_inputs doesn't return reference log probs anymore. We need to calculate it ourselves. # https://github.com/huggingface/trl/blob/05bc43e960396581e458195b8388efe6b82cae1f/trl/trainer/grpo_trainer.py#L1328 if self.beta != 0.0: + print("!!!!!!!!!!") + print("!!!!!!!!!!") with torch.inference_mode(), model.disable_adapter(): ref_per_token_logps = self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep) else: ref_per_token_logps = None # per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1 - + print("!!!!!!!!!!") + print("!!!!!!!!!!") + print(ref_per_token_logps) # x - x.detach() allows for preserving gradients from x advantages = inputs["advantages"] # per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1)