From cac515ee2e5f992823e952a115024fd0a2610341 Mon Sep 17 00:00:00 2001 From: Charles London <36036324+le-big-mac@users.noreply.github.com> Date: Tue, 25 Feb 2025 23:22:35 +0000 Subject: [PATCH] Fix key error in GRPOTrainer (#1818) * fix keyerror in GRPOTrainer * check for train in _metrics --- unsloth/models/rl_replacements.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 06ae82140b..f88e362beb 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -164,7 +164,7 @@ RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__prepare_inputs) # Remove _move_model_to_vllm def grpo_trainer__move_model_to_vllm(function_name, function): if function_name != "_move_model_to_vllm": return function - + def _move_model_to_vllm(self, *args, **kwargs): return None function = inspect.getsource(_move_model_to_vllm) @@ -246,14 +246,20 @@ def grpo_trainer_compute_loss(function_name, function): self, _input_ids, logits_to_keep, completion_mask, advantages, n_chunks = self.args.unsloth_num_chunks, ) - + # 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()) # mean_kl = ((per_token_kl * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean() # self._metrics["kl"].append(self.accelerator.gather_for_metrics(mean_kl).mean().item()) - self._metrics["kl"].append(mean_kl.item()) + + if "train" in self._metrics: + mode = "eval" if self.control.should_evaluate else "train" + self._metrics[mode]["completion_length"].append(completion_length.item()) + self._metrics[mode]["kl"].append(mean_kl.item()) + else: + self._metrics["completion_length"].append(completion_length.item()) + self._metrics["kl"].append(mean_kl.item()) return loss pass