From b414e43b74369d3be678e1ed88fa47d09b1ffa50 Mon Sep 17 00:00:00 2001 From: Datta Nimmaturi Date: Tue, 9 Dec 2025 13:06:22 +0530 Subject: [PATCH] Remove reload_weights rpc call from grpo trainer (#3673) * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Remove reload_weights rpc call from grpo trainer * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Use regex instead of static string * patch openenv reload_weights call * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Better handle sleep and wakeup * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Reset indentation --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han --- unsloth/models/rl.py | 9 ++++ unsloth/models/rl_replacements.py | 69 +++++++++++++++++++++++++++++++ 2 files changed, 78 insertions(+) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index e76ad98e09..005d3b3c9d 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -32,6 +32,7 @@ from .rl_replacements import ( RL_PRE_ITEMS, RL_CONFIG_CHANGES, RL_METRICS_CHANGES, + RL_ADDITIONAL_FUNCTIONS, ) torch_compile_options = { @@ -1327,9 +1328,17 @@ def patch_trl_rl_trainers(): return +def patch_trl_openenv(): + for function in RL_ADDITIONAL_FUNCTIONS["openenv"]: + print(f"Unsloth: Patching trl openenv with function: {function.__name__}") + function() # Call the function to apply the patch + return + + def PatchFastRL(algorithm = None, FastLanguageModel = None): if FastLanguageModel is not None: PatchRL(FastLanguageModel) patch_trl_rl_trainers() + patch_trl_openenv() if type(algorithm) is str and algorithm.islower(): PatchRLStatistics(algorithm) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 38973fe30e..d022e26be6 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -42,6 +42,7 @@ 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, @@ -216,6 +217,27 @@ def grpo_trainer__prepare_inputs(function_name, 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": @@ -906,3 +928,50 @@ def grpo_trainer_metrics(RLTrainer_source, RLConfig_source): 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. + try: + import trl.experimental.openenv.utils as openenv_utils + import trl.experimental.openenv as openenv + except ImportError as e: + print(f"Unsloth: Failed to import trl openenv: {e}") + 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: + print("Unsloth: Warning - regex did not match, patch may have failed") + return + + # Execute and explicitly assign to module + local_ns = {} + exec(compile(src, "", "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 + print("Unsloth: Patched trl openenv generate_rollout_completions") + + +RL_ADDITIONAL_FUNCTIONS["openenv"].append(openenv_vllm_reload_weights)