Update rl_replacements.py
This commit is contained in:
parent
3461b987fd
commit
d71dfb1d01
1 changed files with 28 additions and 8 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue