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 <danielhanchen@gmail.com>
This commit is contained in:
Edd 2024-11-15 05:07:29 +04:00 committed by GitHub
commit b69fee4a36

View file

@ -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