[Fix] [TRL] load_lora for multi line llm.chat/generate (#3696)
* [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 * Handle multi line self.llm.chat better * Use logger * re-indent * Stricter regex to replace wildcard --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han <danielhanchen@gmail.com>
This commit is contained in:
parent
9e5b4052e5
commit
89787329d3
2 changed files with 15 additions and 8 deletions
|
|
@ -24,6 +24,7 @@ import os
|
|||
import re
|
||||
import torch
|
||||
from unsloth_zoo.compiler import create_new_function
|
||||
from unsloth_zoo.log import logger
|
||||
from unsloth_zoo.logging_utils import PatchRLStatistics
|
||||
from unsloth_zoo.rl_replacements import RL_REPLACEMENTS
|
||||
from .rl_replacements import (
|
||||
|
|
@ -349,12 +350,12 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
and trainer_file.split("_")[0] in x.lower()
|
||||
]
|
||||
if len(name) != 1:
|
||||
print(
|
||||
logger.info(
|
||||
f"Unsloth: Could not find Trainer class in trl.trainer.{trainer_file}. Found: {name}"
|
||||
)
|
||||
return
|
||||
if len(config) != 1:
|
||||
print(
|
||||
logger.info(
|
||||
f"Unsloth: Could not find Config class in trl.trainer.{trainer_file}. Found: {config}"
|
||||
)
|
||||
return
|
||||
|
|
@ -365,14 +366,14 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
try:
|
||||
RLTrainer = eval(f"trl.trainer.{trainer_file}.{RLTrainer_name}")
|
||||
except Exception as e:
|
||||
print(
|
||||
logger.info(
|
||||
f"Unsloth: Could not load {RLTrainer_name} from trl.trainer.{trainer_file}: {e}"
|
||||
)
|
||||
return
|
||||
try:
|
||||
RLConfig = eval(f"trl.trainer.{trainer_file}.{RLConfig_name}")
|
||||
except Exception as e:
|
||||
print(
|
||||
logger.info(
|
||||
f"Unsloth: Could not load {RLConfig_name} from trl.trainer.{trainer_file}: {e}"
|
||||
)
|
||||
return
|
||||
|
|
@ -1269,6 +1270,11 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import
|
|||
+ r", load_tensors = True))",
|
||||
source,
|
||||
)
|
||||
# All these are to fix multiple commas before lora_request (in case the original code ends with something like ",)")
|
||||
# https://github.com/huggingface/trl/blob/main/trl/trainer/grpo_trainer.py#L1388 for eg has such an ending
|
||||
source = re.sub(r"\,[\s]{1,}\,[\s]{0,}lora_request", ", lora_request", source)
|
||||
source = re.sub(r"[\s]{1,}\,[\s]{0,}lora_request", ", lora_request", source)
|
||||
source = re.sub(r"[\,]{1,}[\s]{0,}lora_request", ", lora_request", source)
|
||||
# Prefer using unsloth's sampling params and fallback to trl's if not found
|
||||
# We'll enable this later separately when combining both this and GRPOConfig params
|
||||
# source = re.sub(
|
||||
|
|
@ -1330,7 +1336,7 @@ def patch_trl_rl_trainers():
|
|||
|
||||
def patch_trl_openenv():
|
||||
for function in RL_ADDITIONAL_FUNCTIONS["openenv"]:
|
||||
print(f"Unsloth: Patching trl openenv with function: {function.__name__}")
|
||||
logger.info(f"Unsloth: Patching trl openenv with function: {function.__name__}")
|
||||
function() # Call the function to apply the patch
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ import torch
|
|||
import inspect
|
||||
from collections import defaultdict
|
||||
from unsloth_zoo.rl_replacements import RL_REPLACEMENTS, left_pack_padding
|
||||
from unsloth_zoo.log import logger
|
||||
from ..device_type import (
|
||||
is_hip,
|
||||
get_device_type,
|
||||
|
|
@ -945,7 +946,7 @@ def openenv_vllm_reload_weights():
|
|||
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}")
|
||||
logger.warning(f"Unsloth: Failed to import trl openenv: {e}")
|
||||
return
|
||||
|
||||
src = inspect.getsource(openenv_utils.generate_rollout_completions)
|
||||
|
|
@ -960,7 +961,7 @@ def openenv_vllm_reload_weights():
|
|||
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")
|
||||
logger.warning("Unsloth: Warning - regex did not match, patch may have failed")
|
||||
return
|
||||
|
||||
# Execute and explicitly assign to module
|
||||
|
|
@ -971,7 +972,7 @@ def openenv_vllm_reload_weights():
|
|||
# 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")
|
||||
logger.info("Unsloth: Patched trl openenv generate_rollout_completions")
|
||||
|
||||
|
||||
RL_ADDITIONAL_FUNCTIONS["openenv"].append(openenv_vllm_reload_weights)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue