From b69fee4a363ea7a0fb17ca0b0ecf5fafce4442bf Mon Sep 17 00:00:00 2001 From: Edd <68678137+Erland366@users.noreply.github.com> Date: Fri, 15 Nov 2024 05:07:29 +0400 Subject: [PATCH] fix/sfttrainer-compatibility (#1293) * Refactor trainer.py to import SFTConfig directly and update UnslothTrainingArguments class inheritance * Update trainer.py * Update trainer.py --------- Co-authored-by: Daniel Han --- unsloth/trainer.py | 23 +++++++++++------------ 1 file changed, 11 insertions(+), 12 deletions(-) diff --git a/unsloth/trainer.py b/unsloth/trainer.py index 00956ed41b..14fd1631df 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -20,11 +20,6 @@ from functools import wraps import trl import inspect from trl import SFTTrainer -try: - from trl import SFTConfig as TrainingArguments -except: - from transformers import TrainingArguments -pass from . import is_bfloat16_supported from unsloth_zoo.training_utils import unsloth_train as _unsloth_train from packaging.version import Version @@ -60,7 +55,11 @@ else: pass pass - +try: + from trl import SFTConfig as TrainingArguments +except: + from transformers import TrainingArguments +pass @dataclass class UnslothTrainingArguments(TrainingArguments): embedding_learning_rate : Optional[float] = field( @@ -134,7 +133,7 @@ pass # From `trl>=0.13.0`, they changed how to pass several params to the trainer # We need to patch to make the transition smooth -def create_backwards_compatible_trainer(trainer_class, config_class): +def _backwards_compatible_trainer(trainer_class, config_class): original_init = trainer_class.__init__ @wraps(original_init) @@ -167,6 +166,7 @@ def create_backwards_compatible_trainer(trainer_class, config_class): } # Get parameters that exist in Config but not in TrainingArguments + from transformers import TrainingArguments moved_params = \ set(inspect.signature(config_class) .parameters.keys()) - \ set(inspect.signature(TrainingArguments).parameters.keys()) @@ -207,14 +207,13 @@ def _patch_trl_trainer(): import trl.trainer trl_classes = dir(trl.trainer) - - non_convertable_trainer = set(["PPOv2", "AlignProp"]) - trl_trainers = set(x[:-len("Trainer")] for x in trl_classes if x.endswith("Trainer")) - non_convertable_trainer - trl_configs = set(x[:-len("Config")] for x in trl_classes if x.endswith("Config")) - non_convertable_trainer + trl_trainers = set(x[:-len("Trainer")] for x in trl_classes if x.endswith("Trainer")) + trl_configs = set(x[:-len("Config")] for x in trl_classes if x.endswith("Config")) trl_classes = list(trl_trainers & trl_configs) for x in trl_classes: - exec(f"trl.{x}Trainer.__init__ = create_backwards_compatible_trainer(trl.{x}Trainer, trl.{x}Config)", globals()) + try: exec(f"trl.{x}Trainer.__init__ = _backwards_compatible_trainer(trl.{x}Trainer, trl.{x}Config)", globals()) + except: continue pass trl.__UNSLOTH_BACKWARDS_COMPATIBLE__ = True