Fix key error in GRPOTrainer (#1818)

* fix keyerror in GRPOTrainer

* check for train in _metrics
This commit is contained in:
Charles London 2025-02-25 23:22:35 +00:00 committed by GitHub
commit cac515ee2e

View file

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