From 473554810a03e98967d94fa01f0ba3471f35bd5b Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 31 Dec 2023 18:26:17 +1100 Subject: [PATCH] DPO --- unsloth/models/__init__.py | 2 +- unsloth/models/dpo.py | 45 +++++++++++++++++++++++++------------- 2 files changed, 31 insertions(+), 16 deletions(-) diff --git a/unsloth/models/__init__.py b/unsloth/models/__init__.py index 9322049dc8..b174a2cec1 100644 --- a/unsloth/models/__init__.py +++ b/unsloth/models/__init__.py @@ -12,6 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from .loader import FastLanguageModel +from .loader import FastLanguageModel, FastDPOTrainer from .llama import FastLlamaModel from .mistral import FastMistralModel diff --git a/unsloth/models/dpo.py b/unsloth/models/dpo.py index 59e614983f..55ee247bf6 100644 --- a/unsloth/models/dpo.py +++ b/unsloth/models/dpo.py @@ -17,6 +17,9 @@ from transformers.utils.notebook import ( NotebookTrainingTracker, NotebookProgressCallback, ) +from transformers.trainer import DEFAULT_PROGRESS_CALLBACK +from trl import DPOTrainer +import types DPOTrainer_metrics = [ "rewards/chosen", @@ -46,23 +49,35 @@ def NotebookProgressCallback_on_log(self, args, state, control, logs=None, **kwa if args.evaluation_strategy == IntervalStrategy.NO and "loss" in logs: values = {"Training Loss": logs["loss"]} for metric in DPOTrainer_metrics: - values[metric.replace("/", " / ")] = logs[metric] + if metric in logs: + values[metric.replace("/", " / ")] = logs[metric] + else: + # Maybe not a DPO Trainer anymore? Redo the tracker + column_names = [self.first_column] + ["Training Loss"] + if args.evaluation_strategy != IntervalStrategy.NO: + column_names.append("Validation Loss") + self.training_tracker = NotebookTrainingTracker(state.max_steps, column_names) + break + pass + pass # First column is necessarily Step since we're not in epoch eval strategy values["Step"] = state.global_step self.training_tracker.write_line(values) -pass - - -def patch_dpo_trainer(): - # We patch Jupyter Notebook's printing to include all columns for DPO. - NotebookProgressCallback.on_train_begin = NotebookProgressCallback_on_train_begin - NotebookProgressCallback.on_log = NotebookProgressCallback_on_log -pass -# Patch DPO notebook printing -patch_dpo_trainer() - - -from trl import DPOTrainer -class FastDPOTrainer(DPOTrainer): + pass +pass + + +class FastDPOTrainer(DPOTrainer): + # 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 pass