From 9921fb0ffdf0a45d8199e5e5c3fbef64b2b7389c 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 5a934467887d5a6a180d5f4b2e356991b65ed4b8 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 dbc927fd713de4d2f41e4a4fbd4c443ca5537805 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 09ddc2ed5ca746752c053c78eb0560107915fc4e 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