From 2a26cdc184ecb0ec432e5873a82ba4519593ffa7 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 16 Feb 2025 21:06:47 -0800 Subject: [PATCH] Update rl_replacements.py --- unsloth/models/rl_replacements.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 53ec6e6cde..77d7e6a530 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -216,7 +216,8 @@ def grpo_trainer_compute_loss(function_name, function): # attention_mask = torch.cat([prompt_mask, completion_mask], dim=1) attention_mask = None logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens - + _input_ids = input_ids + _logits_to_keep = logits_to_keep per_token_logps = self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep) # Compute the KL divergence between the model and the reference model @@ -238,8 +239,8 @@ def grpo_trainer_compute_loss(function_name, function): if RL_REPLACEMENTS["count"] == 10: raise else: RL_REPLACEMENTS["count"] = 1 RL_REPLACEMENTS["data"] = ( - ref_per_token_logps, per_token_logps, input_ids, completion_mask, self.beta, advantages, - loss, completion_length, mean_kl, completion_ids, + ref_per_token_logps, per_token_logps, _input_ids, completion_mask, self.beta, advantages, + loss, completion_length, mean_kl, completion_ids, _logits_to_keep, ) # Log the metrics # completion_length = self.accelerator.gather_for_metrics(completion_mask.sum(1)).float().mean().item()