diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index e343a2ec70..6ea8547adc 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -320,16 +320,11 @@ 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) @@ -339,10 +334,13 @@ def grpo_trainer_compute_loss(function_name, function): old_hidden_states = inputs["old_per_token_logps"] else: old_hidden_states = None + input_ids = input_ids[:, -logits_to_keep:] if per_token_logps is not None: - ref_per_token_logps = ref_per_token_logps[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred + if ref_per_token_logps is not None: + ref_per_token_logps = ref_per_token_logps[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred + per_token_logps = per_token_logps[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred loss, completion_length, mean_kl = grpo_compute_loss_slow(