From 35d22b8e92de0cd394043a72f74fda65b4ab364f Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 19 Aug 2024 15:50:16 -0700 Subject: [PATCH] Update tokenizer_utils.py --- unsloth/tokenizer_utils.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/unsloth/tokenizer_utils.py b/unsloth/tokenizer_utils.py index 7316656b2a..873544007d 100644 --- a/unsloth/tokenizer_utils.py +++ b/unsloth/tokenizer_utils.py @@ -1109,6 +1109,7 @@ from inspect import getsource import trl.trainer.sft_trainer from trl.trainer.sft_trainer import * from transformers.trainer import * +from trl.trainer.sft_trainer import neftune_post_forward_hook def patch_sft_trainer_tokenizer(): """ @@ -1173,6 +1174,17 @@ def patch_sft_trainer_tokenizer(): "\n"\ "fix_untrained_tokens(self.model, self.tokenizer, self.train_dataset, eps = 1e-16)\n\n" + # Add NEFTune since it doesn't seem to work?? We need to manually inject it + check_text += \ + "\n\n"\ + "if getattr(self.model.get_input_embeddings(), 'neftune_noise_alpha', None) is not None:\n"\ + " if hasattr(self, 'neftune_hook_handle'):\n"\ + " self.neftune_hook_handle.remove()\n"\ + " if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle\n"\ + "\n"\ + " self.neftune_hook_handle = self.model.get_input_embeddings().register_forward_hook(neftune_post_forward_hook)\n\n"\ + "\n" + check_text = check_text.split("\n") check_text = "\n".join(" "*where + x for x in check_text)