From 4ea5249e65a67a97214dc6dde33343f6bc05d3c7 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Thu, 10 Jul 2025 18:24:53 +0000 Subject: [PATCH] Fix argument mismatch in GRPO _get_per_token_logps lambda function --- unsloth/models/rl_replacements.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 002e4d1863..d6d9352d98 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -176,7 +176,7 @@ def grpo_trainer__prepare_inputs(function_name, function): import re # This matches the function signature, decorators and any comments immediately following pattern = r"(\s*@profiling_decorator\s*\n\s*def _prepare_inputs\s*\([^\)]*\)\s*(->\s*[^:]+)?\s*:\s*\n(?:[ ]*#[^\n]*\n)*)" - + match = re.search(pattern, function) insert = ( " if hasattr(self, 'llm'):\n" @@ -196,7 +196,7 @@ def grpo_trainer__prepare_inputs(function_name, function): rest_of_function, flags=re.DOTALL | re.MULTILINE ) - + # We also need to remove the old wake up call from the beginning of the function # since it's injected before the comments. header_and_comments = re.sub( @@ -373,7 +373,7 @@ def grpo_trainer_compute_loss(function_name, function): _input_ids = input_ids _logits_to_keep = logits_to_keep - get_logps_func = lambda model, input_ids, attention_mask, logits_to_keep, batch_size=None, compute_entropy=False: self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep, batch_size) if hasattr(self, "_get_per_token_logps") else self._get_per_token_logps_and_entropies(model, input_ids, attention_mask, logits_to_keep, batch_size, compute_entropy)['logps'] + get_logps_func = lambda model, input_ids, attention_mask, logits_to_keep, batch_size=None, compute_entropy=False: self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep) if hasattr(self, "_get_per_token_logps") else self._get_per_token_logps_and_entropies(model, input_ids, attention_mask, logits_to_keep, batch_size, compute_entropy)['logps'] per_token_logps = get_logps_func(model, input_ids, attention_mask, logits_to_keep)