unsloth/unsloth/models/rl_replacements.py
Daniel Han ab95425653
Fix correctness bugs in rl.py, rl_replacements.py, and vision.py (#3811)
* Fix correctness bugs in rl.py, rl_replacements.py, and vision.py

1. rl_replacements.py (lines 864, 870): Fixed undefined `nanmin`/`nanmax`
   functions by using `.nan_to_num(nan=inf/-inf).min()/.max()` pattern.
   PyTorch doesn't have torch.nanmin/nanmax, so we replace NaN values
   before computing min/max.

2. vision.py (line 150): Fixed bug where code checked for "input" key
   but then accessed kwargs["input_ids"] instead of kwargs["input"].

3. vision.py (line 159): Fixed bug where literal string "key" was used
   instead of the variable `key` when accessing kwargs.

4. rl.py (lines 903, 905): Fixed non-existent `MathError` exception
   by replacing with `ValueError`.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2025-12-31 21:35:48 -08:00

995 lines
42 KiB
Python

# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
__all__ = [
"RL_EXTRA_ARGS",
"RL_FUNCTIONS",
"RL_PRE_ITEMS",
"RL_CONFIG_CHANGES",
"RL_METRICS_CHANGES",
]
import os
import re
import torch
import inspect
from collections import defaultdict
from unsloth_zoo.rl_replacements import RL_REPLACEMENTS, left_pack_padding
from unsloth_zoo.utils import Version
from importlib.metadata import version as importlib_version
from unsloth_zoo.log import logger
import importlib.util
from ..device_type import (
is_hip,
get_device_type,
DEVICE_TYPE,
DEVICE_TYPE_TORCH,
DEVICE_COUNT,
ALLOW_PREQUANTIZED_MODELS,
)
import textwrap
from ._utils import _get_inference_mode_context_manager
RL_EXTRA_ARGS = defaultdict(list)
RL_FUNCTIONS = defaultdict(list)
RL_PRE_ITEMS = defaultdict(list)
RL_CONFIG_CHANGES = defaultdict(list)
RL_METRICS_CHANGES = defaultdict(list)
RL_ADDITIONAL_FUNCTIONS = defaultdict(list)
torch_compile_options = {
"epilogue_fusion": True,
"max_autotune": True,
"shape_padding": True,
"trace.enabled": False,
"triton.cudagraphs": False,
}
# Check untrained tokens
def sft_trainer_fix_untrained_tokens(call_args, extra_args):
if "model" in call_args and "train_dataset" in call_args:
fix_tokenizer = (
"IGNORED_TOKENIZER_NAMES = os.environ.get('UNSLOTH_IGNORED_TOKENIZER_NAMES', '').split('\\n')\n"
"from unsloth_zoo.tokenizer_utils import fix_untrained_tokens\n"
"from unsloth_zoo.training_utils import fix_zero_training_loss\n"
"if 'tokenizer' not in locals(): tokenizer = processing_class\n"
"fix_untrained_tokens(model, tokenizer, train_dataset, IGNORED_TOKENIZER_NAMES, eps = 1e-16)\n"
"fix_zero_training_loss(model, tokenizer, train_dataset)\n"
)
return fix_tokenizer
return ""
RL_EXTRA_ARGS["sft_trainer"].append(sft_trainer_fix_untrained_tokens)
# Remove DPO columns which might randomnly be tokenized
def dpo_trainer_fix_columns(call_args, extra_args):
if "model" in call_args and "train_dataset" in call_args:
fix_dpo = (
"if hasattr(train_dataset, 'column_names'):\n"
" column_names = set(train_dataset.column_names)\n"
" check = ['chosen', 'rejected', 'prompt', 'chosen_input_ids', 'chosen_attention_mask',\n"
" 'chosen_labels', 'rejected_input_ids', 'rejected_attention_mask', 'rejected_labels',\n"
" 'prompt_input_ids', 'prompt_attention_mask']\n"
" if all(x in column_names for x in check):\n"
" train_dataset = train_dataset.remove_columns(['chosen', 'rejected', 'prompt'])\n"
" del check, column_names\n"
)
return fix_dpo
return ""
RL_EXTRA_ARGS["dpo_trainer"].append(dpo_trainer_fix_columns)
# Fix tokenizer double BOS
def sft_trainer_prepare_dataset(function_name, function):
if (
function_name != "_prepare_non_packed_dataloader"
and function_name != "_prepare_dataset"
):
return function
fast_sft_prepare_dataset = RL_REPLACEMENTS.get("sft_prepare_dataset", None)
if fast_sft_prepare_dataset is not None:
params = inspect.signature(fast_sft_prepare_dataset).parameters.keys()
params = ".*?".join(params)
matched = re.match(
r"[\s]{0,}def _prepare_dataset\(.*?" + params + r".*?\)",
function,
flags = re.MULTILINE | re.DOTALL,
)
if matched:
# Use fast version!
function = inspect.getsource(fast_sft_prepare_dataset)
function = function.split("\n")
function = "\n".join(" " * 4 + x for x in function)
function = function.replace(
"def sft_prepare_dataset", "def _prepare_dataset"
)
return function
check_text = (
"if 'skip_prepare_dataset' in locals() and skip_prepare_dataset:\n"
" return dataset\n"
"if 'tokenizer' not in locals(): tokenizer = processing_class\n"
"if 'formatting_func' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `formatting_func` does not exist!')\n"
"if 'dataset_text_field' not in locals() and 'args' in locals(): dataset_text_field = args.dataset_text_field\n"
"if 'dataset_text_field' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `dataset_text_field` does not exist!')\n"
"test_text = dataset[0][dataset_text_field] if (formatting_func is None and dataset_text_field is not None) else formatting_func(dataset[0])[0]\n"
"chat_template = getattr(tokenizer, 'chat_template', None)\n"
"chat_template = '' if chat_template is None else chat_template\n"
"has_bos_token_already = (test_text.startswith(tokenizer.bos_token) or tokenizer.bos_token in chat_template) "
"if getattr(tokenizer, 'bos_token', None) is not None else False\n"
"if 'add_special_tokens' not in locals() and has_bos_token_already:\n"
" from functools import partial\n"
" tokenizer_call = tokenizer.__call__\n"
" tokenizer.__call__ = partial(tokenizer_call, add_special_tokens = False)\n"
" processing_class = tokenizer\n"
"else:\n"
" tokenizer_call = None\n"
" add_special_tokens = False if has_bos_token_already else locals().get('add_special_tokens', False)\n"
)
check_text = check_text.split("\n")
check_text = "\n".join(" " * 8 + x for x in check_text)
check_text = check_text.rstrip() + "\n"
# .*? matches first match. .+? matches final match.
replacer = re.findall(
r"def " + function_name + r"\(.*?\).*?\:\n",
function,
flags = re.MULTILINE | re.DOTALL,
)
if len(replacer) != 0:
replacer = replacer[0]
function = function.replace(replacer, replacer + check_text)
# Return tokenizer's original state
return_state = (
"if tokenizer_call is not None: tokenizer.__call__ = tokenizer_call\n"
)
function = re.sub(
r"\n([ ]{4,})(return .*?[\s]{0,})$",
rf"\1{return_state}\1\2",
function,
)
return function
RL_FUNCTIONS["sft_trainer"].append(sft_trainer_prepare_dataset)
# Ignore mean_token_accuracy since it needs logits
# We override it directly with our version
def sft_trainer_compute_loss(function_name, function):
if function_name != "compute_loss":
return function
def compute_loss(
self, model, inputs, return_outputs = False, num_items_in_batch = None
):
outputs = super().compute_loss(
model,
inputs,
return_outputs = return_outputs,
num_items_in_batch = num_items_in_batch,
)
return outputs
function = inspect.getsource(compute_loss)
return function
RL_FUNCTIONS["sft_trainer"].append(sft_trainer_compute_loss)
# Autocast precision for GRPO
def grpo_trainer__prepare_inputs(function_name, function):
if function_name != "_prepare_inputs":
return function
# Add mixed precision training
function = function.replace(
"with torch.inference_mode():",
"with torch.inference_mode(), "
"torch.amp.autocast(device_type = 'cuda', "
"dtype = ((torch.float16 if os.environ.get('ACCELERATE_MIXED_PRECISION', 'fp16') == 'fp16' else torch.bfloat16) "
"if not torch.is_autocast_enabled('cuda') else nullcontext())"
"if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '0' else torch.float16):",
)
function = function.replace(
"self.accelerator.unwrap_model(self.model)",
"self.accelerator.unwrap_model(self.model, keep_fp32_wrapper = False)",
)
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__prepare_inputs)
# Remove collective RPC of reload weights from generate
# trl added reload weights (potentially for quantized models), we don't need it for our use case (LoRA primarily)
# https://github.com/huggingface/trl/commit/7856d3b1f6518601732f489883b341bb6dd36434#diff-964e6fd373aa93037604064cb2b822d7f8e2735e33f791065acf2c4c3552d393R1168-R1169
def grpo_trainer__generate_single_turn(function_name, function):
if function_name != "_generate_single_turn":
return function
# Remove the reload_weights collective RPC call from the generate function's source
# function = function.replace('self.llm.collective_rpc("reload_weights")', "")
# The regex below does the same thing but is more flexible and can handle single or double quotes
function = re.sub(
r"self\.llm\.collective_rpc\(\s*(['\"])reload_weights\1\s*\)",
"",
function,
)
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__generate_single_turn)
# Fix incorrect special tokens handling and truncation in older TRL versions
def grpo_trainer__generate_and_score_completions(function_name, function):
if function_name != "_generate_and_score_completions":
return function
# TRL 0.19.0 did skip_special_tokens = True which should be False
function = function.replace(
"prompt_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False",
"prompt_ids, skip_special_tokens=False, clean_up_tokenization_spaces=False",
)
# Left pad prompt before calculation old and ref hidden states
line_to_replace = 'batch_size = self.args.per_device_train_batch_size if mode == "train" else self.args.per_device_eval_batch_size'
# The new multi-line string that will replace the line above
replacement_lines = """
batch_size = self.args.per_device_train_batch_size if mode == "train" else self.args.per_device_eval_batch_size
try:
# TRL 0.23.1 and below path
if not has_images:
# Left pad prompt before calculation old and ref hidden states
prompt_completion_ids = left_pack_padding(prompt_completion_ids, self.processing_class.pad_token_id)
self.model.for_training()
except:
# TRL 0.24.0 and below path
if images is None:
# Left pad prompt before calculation old and ref hidden states
prompt_completion_ids = left_pack_padding(prompt_completion_ids, self.processing_class.pad_token_id)
self.model.for_training()"""
function = function.replace(line_to_replace, replacement_lines)
pattern_to_find = re.compile(
r"^\s*if self\.args\.gradient_accumulation_steps % generate_every != 0 or \(\s*"
r"self\.use_vllm and self\.vllm_importance_sampling_correction\s*"
r"\):",
re.MULTILINE,
)
replacement_text = """
if self.args.gradient_accumulation_steps % generate_every != 0 or (
self.use_vllm
):"""
# Use re.sub() to perform the replacement
function, num_replacements = pattern_to_find.subn(replacement_text, function)
pattern_to_find = re.compile(
r"(^\s*)all_logprobs = \[" # Capture indentation (group 1)
r".*?" # Match everything inside non-greedily
r"for output in outputs\.outputs\s*"
r"\]",
re.DOTALL | re.MULTILINE,
)
replacement_text = (
r"\1from trl.scripts.vllm_serve import sanitize_logprob\n"
r"\1all_logprobs = [\n"
r"\1 [sanitize_logprob(next(iter(logprob.values()))) for logprob in output.logprobs]\n"
r"\1 for outputs in all_outputs\n"
r"\1 for output in outputs.outputs\n"
r"\1]"
)
function, num_replacements = pattern_to_find.subn(replacement_text, function)
# Always between max_prompt_length and use_vllm
found = re.findall(
r"\n(([ ]{8,})if self\.max_prompt_length is not None:.*?"
r"\2if self\.use_vllm:)",
function,
flags = re.DOTALL | re.MULTILINE,
)
if len(found) != 0:
replace_part, spacing = found[0]
removed_comments = re.sub(r"\#[^\n]{1,}", "", replace_part)
splits = removed_comments.split("\n")
if (
sum(re.match(rf"{spacing}[^\s]", x) is not None for x in splits) == 2
and len(spacing) >= 8
):
new_replacement = f"""\n{spacing}if self.max_prompt_length is not None:
# If max_prompt_length is set, we trim the prompt to keep only the last `max_prompt_length` tokens.
# Then we decode those tokens back into text. We manually remove leading pad tokens from the decoded text,
# because we can't use `skip_special_tokens=True` (some special tokens are still needed for generation).
protected = [self.image_token_id, self.vision_start_token_id, self.vision_end_token_id]
protected = [token for token in protected if token is not None]
prompt_ids, prompt_mask = truncate_with_protected_tokens(
prompt_ids, prompt_mask, self.max_prompt_length, protected
)
prompts_text = [re.sub(rf"^({{re.escape(self.pad_token)}})+", "", text) for text in prompts_text]
# The chat template inserts a single image token into the prompt text. However, when this text is later
# tokenized, the single image token string is expanded into multiple image token IDs, depending on the
# image size. Since we're detokenizing here, we may see repeated image tokens in the decoded text. We
# collapse them back into a single token string to match the original template.
if self.image_token is not None:
prompts_text = [
re.sub(rf"({{re.escape(self.image_token)}})+", self.image_token, text) for text in prompts_text
]
# Generate completions using either vLLM or regular generation
if self.use_vllm:"""
function = function.replace(replace_part, new_replacement)
string_to_find = """ if "image_sizes" in prompt_inputs:
output["image_sizes"] = prompt_inputs["image_sizes"]"""
replacement_string = """ if "image_sizes" in prompt_inputs:
output["image_sizes"] = prompt_inputs["image_sizes"]
if self.use_vllm:
try:
output["sampling_per_token_logps"] = sampling_per_token_logps
except NameError:
output["sampling_per_token_logps"] = None"""
function = function.replace(string_to_find, replacement_string)
if "wake_up()" not in function:
# Sleep functionality has been added to trl in v0.23.0. We do not want to redo this.
# https://github.com/huggingface/trl/commit/edbe8234bc7e528f72ac76607de9d3e4753e2709
pattern = re.compile(r".*self\.llm\.generate\(.*\).*", re.MULTILINE)
matches = list(pattern.finditer(function))
patched = function
# Generally there's only one match. But this is just to make sure we don't miss any.
for match in reversed(matches):
line = match.group(0)
indent_match = re.match(r"(\s*)", line)
indent = indent_match.group(1) if indent_match else ""
wrapped = (
f"{indent}if hasattr(self, 'llm'):\n"
f"{indent} if getattr(self.llm.llm_engine.vllm_config.model_config, 'enable_sleep_mode', False):\n"
f"{indent} self.llm.wake_up()\n"
f"{line}\n\n"
f"{indent}if hasattr(self, 'llm'):\n"
f"{indent} if getattr(self.llm.llm_engine.vllm_config.model_config, 'enable_sleep_mode', False):\n"
f"{indent} self.llm.sleep(os.environ.get('VLLM_SLEEP_MODE', 1))\n"
)
patched = patched[: match.start()] + wrapped + patched[match.end() :]
function = patched
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__generate_and_score_completions)
# Fix {"reasoning_effort" : "high"} not applied
def grpo_trainer_fix_maybe_apply_chat_template(function_name, function):
spaces = function.find("def ")
if spaces % 4 != 0:
return function
spaces += 4
replacement = """
_chat_template_ = getattr(self.processing_class, "chat_template", None)
if _chat_template_ is None: _chat_template_ = ""
_supported_keys_ = set(("prompt", "chosen", "rejected", "completion", "messages", "label"))
prompts_text = []
for _example_ in __INPUTS__REPLACEMENT__:
_tokenizer_kwargs_ = {}
if type(_example_) is not dict:
_example_ = {"prompt": _example_}
_left_keys_ = _example_.keys() - _supported_keys_
for k in _left_keys_:
if k in _chat_template_:
v = _example_[k]
if type(v) is str:
_tokenizer_kwargs_[k] = v
_x_ = maybe_apply_chat_template(_example_, self.processing_class, **_tokenizer_kwargs_)["prompt"]
prompts_text.append(_x_)
"""
replacement = textwrap.dedent(replacement).strip()
replacement = textwrap.indent(replacement, spaces * " ")
replacement = f"\n{replacement}\n"
what = 'prompts_text = [maybe_apply_chat_template(example, self.processing_class)["prompt"] for example in inputs]'
function = function.replace(
what, replacement.replace("__INPUTS__REPLACEMENT__", "inputs")
)
"""prompts_text = [
maybe_apply_chat_template({"prompt": prompt}, self.processing_class)["prompt"] for prompt in prompts
]"""
function = re.sub(
r"prompts_text = \["
r"[\s]{0,}"
r"maybe_apply_chat_template\(\{[\"\']prompt[\"\'][\s]{0,}\:[\s]{0,}prompt[\s]{0,}\}[\s]{0,}\,[\s]{0,}self\.processing_class\)"
r"\[[\"\']prompt[\"\']\] for prompt in prompts"
r"[\s]{0,}"
r"\]",
replacement.replace("__INPUTS__REPLACEMENT__", "prompts"),
function,
)
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer_fix_maybe_apply_chat_template)
# 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)
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__move_model_to_vllm)
# Edit _get_per_token_logps to handle mixed precision
def grpo_trainer__get_per_token_logps(function_name, function):
if function_name != "_get_per_token_logps":
return function
def _get_per_token_logps(
self, model, input_ids, attention_mask, logits_to_keep, compute_efficient = False
):
if True: # os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0':
return None # 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
)
if os.environ.get("UNSLOTH_FORCE_FLOAT32", "0") == "1":
self._autocast_dtype = torch.float16
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
with torch.amp.autocast(device_type = DEVICE_TYPE, dtype = self._autocast_dtype):
# We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded
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 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
# logits = logits[:, -logits_to_keep:]
# return logits
# See https://huggingface.co/blog/the_n_implementation_details_of_rlhf_with_ppo#policy-training-implementation-details
# logits = logits / self.temperature
# logps = selective_log_softmax(logits, input_ids)
# row_indices, col_indices = torch.where(logps < -20)
# # Method 1: Check if tensors have elements
# if len(row_indices) > 0 and len(col_indices) > 0:
# breakpoint() # Breakpoint triggered here
# print("Found high values!")
# return logps # compute logprobs for the input tokens
function = inspect.getsource(_get_per_token_logps)
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__get_per_token_logps)
def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
if function_name != "_get_per_token_logps_and_entropies":
return 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,
compute_efficient = False,
*args,
**kwargs,
):
# if True: # os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0':
# return None, None # logps, entropies Unsloth efficient GRPO
if compute_efficient:
return None, None
else:
# 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
)
if os.environ.get("UNSLOTH_FORCE_FLOAT32", "0") == "1":
self._autocast_dtype = torch.float16
pixel_values, image_grid_thw = (
kwargs.get("pixel_values", None),
kwargs.get("image_grid_thw", None),
)
pixel_attention_mask, image_sizes = (
kwargs.get("pixel_attention_mask", None),
kwargs.get("image_sizes", None),
)
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
unwrapped_model = self.accelerator.unwrap_model(
model, keep_fp32_wrapper = False
)
with torch.amp.autocast(device_type = "cuda", dtype = self._autocast_dtype):
with _get_inference_mode_context_manager(model):
if pixel_values is None:
attention_mask = input_ids != self.processing_class.pad_token_id
attention_mask = attention_mask.to(attention_mask.dtype)
# We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded
logits = unwrapped_model(
input_ids = input_ids,
attention_mask = attention_mask,
pixel_values = pixel_values,
image_grid_thw = image_grid_thw,
pixel_attention_mask = pixel_attention_mask,
image_sizes = image_sizes,
# logits_to_keep = logits_to_keep + 1,
).logits
else:
logits = unwrapped_model(
input_ids = input_ids,
attention_mask = attention_mask,
pixel_values = pixel_values,
image_grid_thw = image_grid_thw,
pixel_attention_mask = pixel_attention_mask,
image_sizes = image_sizes,
logits_to_keep = logits_to_keep + 1,
).logits
entropies = None
if compute_entropy:
from trl.trainer.utils import entropy_from_logits
entropies = entropy_from_logits(logits)
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "0"
# logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
return logits.detach(), 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
# logits = logits[:, -logits_to_keep:]
# return logits
# See https://huggingface.co/blog/the_n_implementation_details_of_rlhf_with_ppo#policy-training-implementation-details
# logits = logits / self.temperature
# logps = selective_log_softmax(logits, input_ids)
# row_indices, col_indices = torch.where(logps < -20)
# # Method 1: Check if tensors have elements
# if len(row_indices) > 0 and len(col_indices) > 0:
# breakpoint() # Breakpoint triggered here
# print("Found high values!")
# return logps # compute logprobs for the input tokens
function = inspect.getsource(_get_per_token_logps_and_entropies)
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__get_per_token_logps_and_entropies)
grpo_compute_loss = RL_REPLACEMENTS["grpo_compute_loss"]
grpo_compute_loss_slow = RL_REPLACEMENTS["grpo_compute_loss_slow"]
UnslothEfficientGRPO = RL_REPLACEMENTS["UnslothEfficientGRPO"]
grpo_accumulated_loss = RL_REPLACEMENTS["grpo_accumulated_loss"]
grpo_update_SamplingParams = RL_REPLACEMENTS["grpo_update_SamplingParams"]
RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(grpo_compute_loss))
RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(UnslothEfficientGRPO))
RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(grpo_accumulated_loss))
RL_PRE_ITEMS["grpo_trainer"].append(grpo_compute_loss_slow)
RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(grpo_update_SamplingParams))
RL_PRE_ITEMS["grpo_trainer"].append(
inspect.getsource(_get_inference_mode_context_manager)
)
# Edit _get_per_token_logps to handle mixed precision
def grpo_trainer_compute_loss(function_name, function):
if function_name != "compute_loss":
return function
def compute_loss(
self, model, inputs, return_outputs = False, num_items_in_batch = None
):
if return_outputs:
raise ValueError("The GRPOTrainer does not support returning outputs")
# Compute the per-token log probabilities for the model
prompt_ids, prompt_mask = inputs["prompt_ids"], inputs["prompt_mask"]
completion_ids, completion_mask = (
inputs["completion_ids"],
inputs["completion_mask"],
)
pixel_values, image_grid_thw = (
inputs.get("pixel_values", None),
inputs.get("image_grid_thw", None),
)
pixel_attention_mask, image_sizes = (
inputs.get("pixel_attention_mask", None),
inputs.get("image_sizes", None),
)
num_items_in_batch = inputs.get("num_items_in_batch", None)
sampling_per_token_logps = inputs.get("sampling_per_token_logps", None)
current_gradient_accumulation_steps = self.current_gradient_accumulation_steps
num_processes = self.accelerator.num_processes
input_ids = torch.cat([prompt_ids, completion_ids], dim = 1)
bsz, qlen = input_ids.shape
attention_mask = torch.cat([prompt_mask, completion_mask], dim = 1)
# attention_mask = None
logits_to_keep = completion_ids.size(
1
) # we only need to compute the logits for the completion tokens
_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,
compute_efficient = False: self._get_per_token_logps(
model, input_ids, attention_mask, logits_to_keep, compute_efficient
)
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,
compute_efficient,
)[0]
) # logps
per_token_logps = get_logps_func(
model, input_ids, attention_mask, logits_to_keep, compute_efficient = True
)
# Compute the KL divergence between the model and the reference model
# _prepare_inputs doesn't return reference log probs anymore. We need to calculate it ourselves.
# https://github.com/huggingface/trl/blob/05bc43e960396581e458195b8388efe6b82cae1f/trl/trainer/grpo_trainer.py#L1328
# if self.beta != 0.0:
# with torch.inference_mode(), model.disable_adapter():
# ref_per_token_logps = per_token_logps = get_logps_func(model, input_ids, attention_mask, logits_to_keep)
# else:
# ref_per_token_logps = None
ref_hidden_states = inputs.get("ref_per_token_logps", None)
# per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1
# x - x.detach() allows for preserving gradients from x
advantages = inputs["advantages"]
# per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1)
# per_token_loss = -(per_token_loss - self.beta * per_token_kl)
# loss = ((per_token_loss * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean()
old_hidden_states = inputs.get("old_per_token_logps", None)
input_ids = input_ids[:, -logits_to_keep:]
# Get logit softcapping and logit scale
logit_softcapping = getattr(model.config, "final_logit_softcapping", 0) # Gemma
if logit_softcapping is None:
logit_softcapping = 0
logit_scale_multiply = getattr(model.config, "logit_scale", 0) # Cohere
if logit_scale_multiply is None:
logit_scale_multiply = 0
logit_scale_divide = getattr(model.config, "logits_scaling", 0) # Granite
if logit_scale_divide is None:
logit_scale_divide = 0
if per_token_logps is not None:
if ref_hidden_states is not None:
ref_hidden_states = ref_hidden_states[
:, :-1, :
] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
if old_hidden_states is not None:
old_hidden_states = old_hidden_states[
:, :-1, :
] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
per_token_logps = per_token_logps[
:, :-1, :
] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
loss, completion_length, mean_kl, delta, flat_is_ratio = (
grpo_compute_loss_slow(
ref_hidden_states,
per_token_logps,
old_hidden_states,
input_ids,
completion_mask,
self.beta,
advantages,
pixel_values = pixel_values,
image_grid_thw = image_grid_thw,
loss_type = self.args.loss_type,
importance_sampling_level = self.importance_sampling_level,
epsilon_low = self.epsilon_low,
epsilon_high = self.epsilon_high,
max_completion_length = self.args.max_completion_length,
delta = self.args.delta,
temperature = self.args.temperature,
logit_softcapping = logit_softcapping,
logit_scale_multiply = logit_scale_multiply,
logit_scale_divide = logit_scale_divide,
num_items_in_batch = num_items_in_batch,
current_gradient_accumulation_steps = current_gradient_accumulation_steps,
num_processes = num_processes,
sampling_per_token_logps = sampling_per_token_logps,
)
)
else:
if hasattr(self.args, "loss_type"):
loss, completion_length, mean_kl, delta, flat_is_ratio = (
grpo_accumulated_loss(
trainer = self,
input_ids = _input_ids,
pixel_values = pixel_values,
image_grid_thw = image_grid_thw,
logits_to_keep = logits_to_keep,
completion_mask = completion_mask,
advantages = advantages,
old_hidden_states = old_hidden_states,
ref_hidden_states = ref_hidden_states,
n_chunks = self.args.unsloth_num_chunks,
loss_type = self.args.loss_type,
importance_sampling_level = self.importance_sampling_level,
epsilon_low = self.epsilon_low,
epsilon_high = self.epsilon_high,
max_completion_length = self.args.max_completion_length,
delta = self.args.delta,
temperature = self.args.temperature,
logit_softcapping = logit_softcapping,
logit_scale_multiply = logit_scale_multiply,
logit_scale_divide = logit_scale_divide,
attention_mask = attention_mask,
num_items_in_batch = num_items_in_batch,
current_gradient_accumulation_steps = current_gradient_accumulation_steps,
num_processes = num_processes,
sampling_per_token_logps = sampling_per_token_logps,
)
)
else:
# to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17
loss, completion_length, mean_kl = grpo_accumulated_loss(
trainer = self,
input_ids = _input_ids,
logits_to_keep = logits_to_keep,
completion_mask = completion_mask,
advantages = advantages,
old_hidden_states = old_hidden_states,
ref_hidden_states = ref_hidden_states,
n_chunks = self.args.unsloth_num_chunks,
temperature = self.args.temperature,
logit_softcapping = logit_softcapping,
logit_scale_multiply = logit_scale_multiply,
logit_scale_divide = logit_scale_divide,
attention_mask = attention_mask,
)
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())
if self.use_vllm and delta is not None:
mean_delta = (
torch.mean(delta)
if delta.numel() > 0
else torch.tensor(0.0, device = self.model.device)
)
max_delta = (
torch.max(delta)
if delta.numel() > 0
else torch.tensor(0.0, device = self.model.device)
)
self._metrics[mode]["sampling/sampling_logp_difference/mean"].append(
self.accelerator.gather(mean_delta).mean().item()
)
self._metrics[mode]["sampling/sampling_logp_difference/max"].append(
self.accelerator.gather(max_delta).max().item()
)
min_importance_sampling_ratio = (
torch.min(flat_is_ratio)
if flat_is_ratio.numel() > 0
else torch.tensor(0.0, device = self.model.device)
)
mean_importance_sampling_ratio = (
torch.mean(flat_is_ratio)
if flat_is_ratio.numel() > 0
else torch.tensor(0.0, device = self.model.device)
)
max_importance_sampling_ratio = (
torch.max(flat_is_ratio)
if flat_is_ratio.numel() > 0
else torch.tensor(0.0, device = self.model.device)
)
self._metrics[mode]["sampling/importance_sampling_ratio/min"].append(
self.accelerator.gather(min_importance_sampling_ratio)
.nan_to_num(nan = float("inf"))
.min()
.item()
)
self._metrics[mode]["sampling/importance_sampling_ratio/mean"].append(
self.accelerator.gather(mean_importance_sampling_ratio).nanmean().item()
)
self._metrics[mode]["sampling/importance_sampling_ratio/max"].append(
self.accelerator.gather(max_importance_sampling_ratio)
.nan_to_num(nan = float("-inf"))
.max()
.item()
)
return loss
function = inspect.getsource(compute_loss)
return function
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer_compute_loss)
# https://github.com/huggingface/trl/blob/main/trl/trainer/grpo_trainer.py#L356
# TRL warns if batch size is not a multiple of num_generations -> fix this.
def grpo_trainer_fix_batch_size(RLTrainer_source, RLConfig_source):
if "divisible by the number of generations" not in RLTrainer_source:
# in later trl versions this doesn't exist anymore
return ""
if "num_generations" not in RLConfig_source:
return ""
check_batch_size = (
"div = per_device_train_batch_size // num_generations\n"
"if div * num_generations != per_device_train_batch_size:\n"
" print('Unsloth: We now expect `per_device_train_batch_size` to be a multiple of `num_generations`.\\n"
"We will change the batch size of ' + str(per_device_train_batch_size) + ' to the `num_generations` of ' + str(num_generations))\n"
" per_device_train_batch_size = num_generations\n"
)
return check_batch_size
RL_CONFIG_CHANGES["grpo_trainer"].append(grpo_trainer_fix_batch_size)
# Add other reward function names
def grpo_trainer_metrics(RLTrainer_source, RLConfig_source):
if "reward_funcs" not in RLTrainer_source:
return ""
# For new TRL we have /mean and /std
use_mean = "rewards/{reward_func_name}/mean" in RLTrainer_source
use_std = "rewards/{reward_func_name}/std" in RLTrainer_source
if not use_mean:
use_normal = "rewards/{reward_func_name}" in RLTrainer_source
else:
use_normal = False
log_metrics = (
"if not isinstance(reward_funcs, list): _reward_funcs = [reward_funcs]\n"
"else: _reward_funcs = reward_funcs\n"
"for reward_func in _reward_funcs:\n"
" try:\n"
" reward_func_name = reward_func.__name__\n"
f" if {use_mean}:\n"
" other_metrics.append(f'rewards/{reward_func_name}/mean')\n"
f" if {use_std}:\n"
" other_metrics.append(f'rewards/{reward_func_name}/std')\n"
f" if {use_normal}:\n"
" other_metrics.append(f'rewards/{reward_func_name}')\n"
" except: pass\n"
)
return log_metrics
RL_METRICS_CHANGES["grpo_trainer"].append(grpo_trainer_metrics)
def openenv_vllm_reload_weights():
# This function patches the trl openenv generate_rollout_completions function to:
# 1. Remove the reload_weights call (unsloth handles weight reloading)
# 2. Fix wake_up call to be compatible with unsloth (remove tags to wake everything)
#
# The issue: TRL's wake_up(tags=["kv_cache"]) only wakes kv_cache, leaving is_sleeping=True
# at the executor level. This causes unsloth's patched generate to try waking up again,
# resulting in double create_and_map on already-mapped handles.
#
# The fix: Use wake_up() with no tags, which wakes everything. Unsloth's patched
# CuMemAllocator.wake_up skips weights anyway, so this is safe.
if importlib.util.find_spec("trl") is None:
return
if Version(importlib_version("trl")) < Version("0.26.0"):
return
try:
import trl.experimental.openenv.utils as openenv_utils
import trl.experimental.openenv as openenv
except ImportError as e:
logger.info(f"Unsloth: Failed to import trl openenv: {e}")
logger.info(
"Unsloth: trl.experimental.openenv not available — skipping RL openenv patches."
)
return
src = inspect.getsource(openenv_utils.generate_rollout_completions)
src = textwrap.dedent(src)
original_src = src
# Remove the reload_weights call - unsloth handles this differently
src = re.sub(r'.*\.collective_rpc\("reload_weights"\).*\n?', "", src)
# Change wake_up(tags=["kv_cache"]) to wake_up() - wake everything to set is_sleeping=False
# This prevents double wake_up issues. Unsloth's allocator skips weights anyway.
src = re.sub(r"\.wake_up\(tags=\[.*?\]\)", ".wake_up()", src)
if original_src == src:
logger.warning("Unsloth: Warning - regex did not match, patch may have failed")
return
# Execute and explicitly assign to module
local_ns = {}
exec(compile(src, "<unsloth>", "exec"), openenv_utils.__dict__, local_ns)
patched_func = local_ns["generate_rollout_completions"]
# Patch both the utils module and the parent openenv module
openenv_utils.generate_rollout_completions = patched_func
openenv.generate_rollout_completions = patched_func
logger.info("Unsloth: Patched trl openenv generate_rollout_completions")
RL_ADDITIONAL_FUNCTIONS["openenv"].append(openenv_vllm_reload_weights)