From a2fa52f28c36df674d7ad46d3572387e14852393 Mon Sep 17 00:00:00 2001 From: Erland366 Date: Thu, 3 Jul 2025 14:08:04 +0000 Subject: [PATCH 1/4] Always use SFTConfig --- unsloth/trainer.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/unsloth/trainer.py b/unsloth/trainer.py index 75fdd410e1..a360205432 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -61,13 +61,9 @@ else: pass pass -try: - from trl import SFTConfig as TrainingArguments -except: - from transformers import TrainingArguments -pass +from trl import SFTConfig @dataclass -class UnslothTrainingArguments(TrainingArguments): +class UnslothTrainingArguments(SFTConfig): embedding_learning_rate : Optional[float] = field( default = None, metadata = {"help" : "Different learning rates for embeddings and lm_head."} From 18c73a42fab60eae1de1d7f806a6218235a8c7bf Mon Sep 17 00:00:00 2001 From: Erland366 Date: Thu, 3 Jul 2025 15:39:56 +0000 Subject: [PATCH 2/4] Refactor UnslothTrainingArguments to support fallback for TrainingArguments import --- unsloth/trainer.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/unsloth/trainer.py b/unsloth/trainer.py index a360205432..e7c73deeaf 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -61,9 +61,13 @@ else: pass pass -from trl import SFTConfig -@dataclass -class UnslothTrainingArguments(SFTConfig): +try: + from trl import SFTConfig as TrainingArguments +except: + from transformers import TrainingArguments +pass + +class UnslothTrainingArguments(TrainingArguments): embedding_learning_rate : Optional[float] = field( default = None, metadata = {"help" : "Different learning rates for embeddings and lm_head."} From fcb0b94f24b935a0cc14cd73e2a0e652690e677b Mon Sep 17 00:00:00 2001 From: Erland366 Date: Thu, 3 Jul 2025 15:47:10 +0000 Subject: [PATCH 3/4] Refactor UnslothTrainingArguments to initialize embedding_learning_rate in constructor --- unsloth/trainer.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/unsloth/trainer.py b/unsloth/trainer.py index e7c73deeaf..3231353296 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -68,10 +68,11 @@ except: pass class UnslothTrainingArguments(TrainingArguments): - embedding_learning_rate : Optional[float] = field( - default = None, - metadata = {"help" : "Different learning rates for embeddings and lm_head."} - ) + def __init__(self, embedding_learning_rate: float = None, *args, **kwargs): + embedding_learning_rate : Optional[float] = field( + default = embedding_learning_rate, + metadata = {"help" : "Different learning rates for embeddings and lm_head."} + ) pass From 3bad31a16ff603fd6c5c1e271df20e5993aa6ec1 Mon Sep 17 00:00:00 2001 From: Erland366 Date: Thu, 3 Jul 2025 15:47:58 +0000 Subject: [PATCH 4/4] Initialize parent class in UnslothTrainingArguments constructor --- unsloth/trainer.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/unsloth/trainer.py b/unsloth/trainer.py index 3231353296..6a42fe4a2f 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -69,10 +69,8 @@ pass class UnslothTrainingArguments(TrainingArguments): def __init__(self, embedding_learning_rate: float = None, *args, **kwargs): - embedding_learning_rate : Optional[float] = field( - default = embedding_learning_rate, - metadata = {"help" : "Different learning rates for embeddings and lm_head."} - ) + embedding_learning_rate = embedding_learning_rate + super().__init__(*args, **kwargs) pass