get_per_token_logps_and_entropies: return tuple instead of dict (#3080)

This commit is contained in:
DoubleMathew 2025-08-02 05:31:41 -05:00 committed by GitHub
commit fbb8a52ce6

View file

@ -352,7 +352,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
# Just copy over from _get_per_token_logps replacement function above. For now this returns None anyway
def _get_per_token_logps_and_entropies(self, model, input_ids, attention_mask, logits_to_keep, batch_size = None, compute_entropy = False, *args, **kwargs):
if True: # os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0':
return {"logps": None, "entropies": None} # Unsloth efficient GRPO
return None, None # logps, entropies Unsloth efficient GRPO
# Otherwise, calculate normally:
if not hasattr(self, '_autocast_dtype'):
self._autocast_dtype = torch.float16 if os.environ.get('ACCELERATE_MIXED_PRECISION', 'fp16') == 'fp16' else torch.bfloat16
@ -373,7 +373,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
entropies = entropy_from_logits(logits)
# logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
return {"logps": logits, "entropies": entropies}
return logits, entropies # logps, entropies
# input_ids = input_ids[:, -logits_to_keep:]
# For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves.
# See https://github.com/huggingface/trl/issues/2770
@ -430,7 +430,7 @@ def grpo_trainer_compute_loss(function_name, function):
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']
self._get_per_token_logps_and_entropies(model, input_ids, attention_mask, logits_to_keep, batch_size, compute_entropy)[0] # logps
per_token_logps = get_logps_func(model, input_ids, attention_mask, logits_to_keep)