Update rl_replacements.py

This commit is contained in:
Daniel Han 2025-06-22 02:55:00 -07:00
commit edb70686d3

View file

@ -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(