From 148e7db6c3115fa1b48706c2af352ece67555fe5 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 31 Dec 2023 18:36:43 +1100 Subject: [PATCH] PatchDPOTrainer --- unsloth/models/__init__.py | 3 ++- unsloth/models/dpo.py | 21 ++++++--------------- unsloth/models/loader.py | 1 - 3 files changed, 8 insertions(+), 17 deletions(-) diff --git a/unsloth/models/__init__.py b/unsloth/models/__init__.py index b174a2cec1..891947d69e 100644 --- a/unsloth/models/__init__.py +++ b/unsloth/models/__init__.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from .loader import FastLanguageModel, FastDPOTrainer +from .loader import FastLanguageModel from .llama import FastLlamaModel from .mistral import FastMistralModel +from .dpo import PatchDPOTrainer diff --git a/unsloth/models/dpo.py b/unsloth/models/dpo.py index 55ee247bf6..e7724c2d0a 100644 --- a/unsloth/models/dpo.py +++ b/unsloth/models/dpo.py @@ -17,9 +17,6 @@ from transformers.utils.notebook import ( NotebookTrainingTracker, NotebookProgressCallback, ) -from transformers.trainer import DEFAULT_PROGRESS_CALLBACK -from trl import DPOTrainer -import types DPOTrainer_metrics = [ "rewards/chosen", @@ -32,6 +29,7 @@ DPOTrainer_metrics = [ "logits/chosen", ] + def NotebookProgressCallback_on_train_begin(self, args, state, control, **kwargs): self.first_column = "Epoch" if args.evaluation_strategy == IntervalStrategy.EPOCH else "Step" self.training_loss = 0 @@ -67,17 +65,10 @@ def NotebookProgressCallback_on_log(self, args, state, control, logs=None, **kwa pass -class FastDPOTrainer(DPOTrainer): +def PatchDPOTrainer(): # Patch DPO notebook printing - if (DEFAULT_PROGRESS_CALLBACK is NotebookProgressCallback): - - DEFAULT_PROGRESS_CALLBACK.on_train_begin = types.MethodType( - NotebookProgressCallback_on_train_begin, - DEFAULT_PROGRESS_CALLBACK, - ) - DEFAULT_PROGRESS_CALLBACK.on_log = types.MethodType( - NotebookProgressCallback_on_log, - DEFAULT_PROGRESS_CALLBACK, - ) - pass + from transformers.trainer import DEFAULT_PROGRESS_CALLBACK + DEFAULT_PROGRESS_CALLBACK.on_train_begin = NotebookProgressCallback_on_train_begin + DEFAULT_PROGRESS_CALLBACK.on_log = NotebookProgressCallback_on_log pass + diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index baa55f9f00..eb8d4960b9 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -14,7 +14,6 @@ from .llama import FastLlamaModel, logger from .mistral import FastMistralModel -from .dpo import FastDPOTrainer from transformers import AutoConfig from transformers import __version__ as transformers_version