diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 96e256b653..df95f73fc5 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -251,7 +251,7 @@ def grpo_trainer__get_per_token_logps(function_name, function): if function_name != "_get_per_token_logps": return function def _get_per_token_logps(self, model, input_ids, attention_mask, logits_to_keep, calc_logprob_flag = None): - if os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0' and not calc_logprob_flag: + if os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0' and not calc_logprob_flag: return None # Unsloth efficient GRPO # Otherwise, calculate normally: if not hasattr(self, '_autocast_dtype'): @@ -337,29 +337,49 @@ def grpo_trainer_compute_loss(function_name, function): 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, :] - per_token_logps = per_token_logps[:, :-1, :] + + 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( - ref_per_token_logps, per_token_logps, old_hidden_states, input_ids, completion_mask, self.beta, advantages, + ref_per_token_logps, + per_token_logps, + old_hidden_states, + input_ids, + completion_mask, + self.beta, + advantages, loss_type = self.args.loss_type, - epsilon_low = self.epsilon_low, epsilon_high = self.epsilon_high, + epsilon_low = self.epsilon_low, + epsilon_high = self.epsilon_high, max_completion_length = self.args.max_completion_length, delta = self.args.delta, ) else: if hasattr(self.args, "loss_type"): loss, completion_length, mean_kl = grpo_accumulated_loss( - self, _input_ids, logits_to_keep, completion_mask, advantages, old_hidden_states, + self, + _input_ids, + logits_to_keep, + completion_mask, + advantages, + old_hidden_states, n_chunks = self.args.unsloth_num_chunks, loss_type = self.args.loss_type, - epsilon_low = self.epsilon_low, epsilon_high = self.epsilon_high, + epsilon_low = self.epsilon_low, + epsilon_high = self.epsilon_high, max_completion_length = self.args.max_completion_length, delta = self.args.delta, ) else: # to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17 loss, completion_length, mean_kl = grpo_accumulated_loss( - self, _input_ids, logits_to_keep, completion_mask, advantages, old_hidden_states, + self, + _input_ids, + logits_to_keep, + completion_mask, + advantages, + old_hidden_states, n_chunks = self.args.unsloth_num_chunks, )