Update rl_replacements.py
This commit is contained in:
parent
e57cb1cbc4
commit
9741989886
1 changed files with 2 additions and 2 deletions
|
|
@ -261,13 +261,13 @@ def grpo_trainer__get_per_token_logps(function_name, function):
|
|||
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
|
||||
with torch.amp.autocast(device_type = 'cuda', dtype = self._autocast_dtype):
|
||||
# We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded
|
||||
hidden_states = model(
|
||||
logits = model(
|
||||
input_ids = input_ids,
|
||||
attention_mask = attention_mask,
|
||||
logits_to_keep = logits_to_keep + 1,
|
||||
).logits
|
||||
# logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
return hidden_states
|
||||
return logits
|
||||
# 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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue