From c4cc776d1033078a00883aab23ecde26623f508a Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 16 Feb 2025 19:47:39 -0800 Subject: [PATCH] Update rl_replacements.py --- unsloth/models/rl_replacements.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 92b12647cf..b058d0d271 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -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())