Update rl_replacements.py

This commit is contained in:
Daniel Han 2025-02-16 19:47:39 -08:00
commit c4cc776d10

View file

@ -233,11 +233,14 @@ def grpo_trainer_compute_loss(function_name, function):
ref_per_token_logps, per_token_logps, input_ids, completion_mask, self.beta, advantages,
)
from unsloth_zoo.rl_replacements import RL_REPLACEMENTS
if "count" in RL_REPLACEMENTS:
RL_REPLACEMENTS["count"] += 1
if RL_REPLACEMENTS["count"] == 5: 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,
)
raise
# Log the metrics
# completion_length = self.accelerator.gather_for_metrics(completion_mask.sum(1)).float().mean().item()
self._metrics["completion_length"].append(completion_length.item())