diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index e1b40e64ed..4ae0ee2ed8 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -1315,10 +1315,17 @@ def patch_trl_rl_trainers(): _patch_trl_rl_trainers(trainer) return +def patch_trl_openenv(): + from unsloth.models.rl_replacements import RL_ADDITIONAL_FUNCTIONS + 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 8549567f8e..7ab147b428 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, @@ -927,3 +928,34 @@ 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 remove the reload_weights call. + 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 + src = re.sub(r'.*\.collective_rpc\("reload_weights"\).*\n?', '', 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)