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:
parent
786aea6365
commit
b69fee4a36
1 changed files with 11 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue