diff --git a/unsloth/tokenizer_utils.py b/unsloth/tokenizer_utils.py index c445879df7..ca337d84e3 100644 --- a/unsloth/tokenizer_utils.py +++ b/unsloth/tokenizer_utils.py @@ -42,6 +42,7 @@ __all__ = [ "check_tokenizer", "add_new_tokens", "fix_sentencepiece_gguf", + "get_tokenizer_info", ] @@ -276,53 +277,14 @@ def assert_same_tokenization(slow_tokenizer, fast_tokenizer): replacement_char = b"\xc3\xaf\xc2\xbf\xc2\xbd".decode("utf-8") all_special_tokens = [x for x in all_special_tokens if x != replacement_char] - # Check if chat template is enabled! check_chat_template1 = True check_chat_template2 = True check_chat_template3 = True - """ - Weirdly Mistral tokenizers are actually correct?? - Ie below will actually load mistral v1 and v3 incorrectly! - - slow_chat_template = getattr(slow_tokenizer, "chat_template", None) - fast_chat_template = getattr(fast_tokenizer, "chat_template", None) - messages = [ - {"role": "user", "content": " What is 2+2? "}, - {"role": "assistant", "content": " It's 4. "}, - ] - # Check the tokenizer's own chat template - if slow_chat_template is not None and fast_chat_template is not None: - check_chat_template1 = \ - slow_tokenizer.apply_chat_template(messages) == \ - fast_tokenizer.apply_chat_template(messages) - pass - - # Check Mistral chat template without BOS / EOS - slow_tokenizer.chat_template = mistral_template - fast_tokenizer.chat_template = mistral_template - check_chat_template2 = \ - slow_tokenizer.apply_chat_template(messages) == \ - fast_tokenizer.apply_chat_template(messages) - pass - - # Check Llama chat template without BOS / EOS - slow_tokenizer.chat_template = llama_template - fast_tokenizer.chat_template = llama_template - check_chat_template3 = \ - slow_tokenizer.apply_chat_template(messages) == \ - fast_tokenizer.apply_chat_template(messages) - pass - - # Combine them all and revert chat templates - slow_tokenizer.chat_template = slow_chat_template - fast_tokenizer.chat_template = fast_chat_template - """ check_chat_template = ( check_chat_template1 and check_chat_template2 and check_chat_template3 ) - # Try special tokens try: string = ( "\n".join(all_special_tokens) @@ -335,9 +297,6 @@ def assert_same_tokenization(slow_tokenizer, fast_tokenizer): return check_chat_template and check_special_tokens except: - # For eg see https://github.com/unslothai/unsloth/issues/292 - # Sometimes tokenizer has weird tokens, causing a combined tokenization to fail. - # [TODO] We temporarily disable this for CodeLlama tokenizers if slow_tokenizer.__repr__().split("(", 1)[0] in IGNORED_TOKENIZER_CHECKING: return check_chat_template else: @@ -350,17 +309,13 @@ def fix_sentencepiece_tokenizer( token_mapping, temporary_location = "_unsloth_sentencepiece_temp", ): - # From https://github.com/google/sentencepiece/issues/121 - # We need to manually edit the sentencepiece tokenizer! try: from transformers.convert_slow_tokenizer import import_protobuf - sentencepiece_model_pb2 = import_protobuf() except Exception as e: try: import google.protobuf from unsloth_zoo.utils import Version - protobuf_version = Version(google.protobuf.__version__) if protobuf_version > Version("3.20.3"): raise RuntimeError( @@ -368,17 +323,14 @@ def fix_sentencepiece_tokenizer( f"Please downgrade via `pip install --force-reinstall protobuf==3.20.3`" ) except: - # This will only work for older SentencePiece versions <= 3.20.3 from transformers.utils import sentencepiece_model_pb2 if not os.path.exists(temporary_location): os.makedirs(temporary_location) - # Check if tokenizer.model exists if not os.path.isfile(f"{temporary_location}/tokenizer.model"): return new_tokenizer - # First save the old tokenizer old_tokenizer.save_pretrained(temporary_location) tokenizer_file = sentencepiece_model_pb2.ModelProto() @@ -386,21 +338,17 @@ def fix_sentencepiece_tokenizer( open(f"{temporary_location}/tokenizer.model", "rb").read() ) - # Now save the new tokenizer new_tokenizer.save_pretrained(temporary_location) - # Now correct the old tokenizer's .model file for old_token, new_token in token_mapping.items(): ids = old_tokenizer([old_token], add_special_tokens = False).input_ids ids = ids[0] if len(ids) != 1: - # Skip this token! print( f"Skip mapping {old_token} to {new_token} since {new_token} is already in the tokenizer!" ) continue ids = ids[0] - # [TODO] Hack for Starling - try except try: tokenizer_piece = tokenizer_file.pieces[ids] except: @@ -408,13 +356,10 @@ def fix_sentencepiece_tokenizer( assert tokenizer_piece.piece == old_token tokenizer_piece.piece = new_token - # And now write it with open(f"{temporary_location}/tokenizer.model", "wb") as file: file.write(tokenizer_file.SerializeToString()) - # And load it! from transformers import AutoTokenizer - tokenizer = AutoTokenizer.from_pretrained( temporary_location, eos_token = new_tokenizer.eos_token, @@ -424,11 +369,6 @@ def fix_sentencepiece_tokenizer( def fix_sentencepiece_gguf(saved_location): - """ - Fixes sentencepiece tokenizers which did not extend the vocabulary with - user defined tokens. - Inspiration from https://github.com/ggerganov/llama.cpp/blob/master/convert_hf_to_gguf.py - """ from copy import deepcopy from transformers.utils import sentencepiece_model_pb2 import json @@ -442,7 +382,6 @@ def fix_sentencepiece_gguf(saved_location): UNUSED = 5 BYTE = 6 - # Load tokenizer.model tokenizer_file = sentencepiece_model_pb2.ModelProto() if not os.path.isfile(f"{saved_location}/tokenizer.model"): return @@ -451,7 +390,6 @@ def fix_sentencepiece_gguf(saved_location): ) sentence_piece_size = len(tokenizer_file.pieces) - # Load added_tokens_json if not os.path.isfile(f"{saved_location}/added_tokens.json"): return with open(f"{saved_location}/added_tokens.json", "r", encoding = "utf-8") as file: @@ -464,7 +402,6 @@ def fix_sentencepiece_gguf(saved_location): ) new_size = sentence_piece_size + len(added_tokens_json) - # Confirm added_tokens_json is correct added_tokens_ids = np.array(list(added_tokens_json.values())) diff = np.diff(added_tokens_ids) if diff.min() != 1 or diff.max() != 1: @@ -472,7 +409,6 @@ def fix_sentencepiece_gguf(saved_location): if added_tokens_ids.min() != sentence_piece_size: return - # Edit sentence piece tokens with added_tokens_json logger.warning( f"Unsloth: Extending {saved_location}/tokenizer.model with added_tokens.json.\n" f"Originally tokenizer.model is of size ({sentence_piece_size}).\n" @@ -488,10 +424,6 @@ def fix_sentencepiece_gguf(saved_location): with open(f"{saved_location}/tokenizer.model", "wb") as file: file.write(tokenizer_file.SerializeToString()) - - # Add padding tokens - # actual_vocab_size = model.config.vocab_size - # padding = actual_vocab_size - len(tokenizer_file.pieces) return @@ -507,14 +439,10 @@ def _load_correct_tokenizer( if IS_COLAB_ENVIRONMENT: cache_dir = cache_dir elif IS_KAGGLE_ENVIRONMENT: - # /tmp of Kaggle seems has a 80GB limit! - # Let's utilize them cache_dir = os.path.join(KAGGLE_TMP, cache_dir) else: cache_dir = None - # Try loading the slow tokenizer. If it fails, then try Fast only - # Mainly to solve Deepseek models with no tokenizer.model file slow_tokenizer = None try: slow_tokenizer = AutoTokenizer.from_pretrained( @@ -523,7 +451,6 @@ def _load_correct_tokenizer( padding_side = padding_side, token = token, trust_remote_code = trust_remote_code, - # Cannot just use use_fast = False as per https://twitter.com/danielhanchen/status/1789659394302718373 use_fast = False, legacy = False, from_slow = True, @@ -531,11 +458,7 @@ def _load_correct_tokenizer( ) except: slow_tokenizer = None - # print( - # f"Unsloth: {tokenizer_name} has no tokenizer.model file.\n"\ - # "Just informing you about this - this is not a critical error." - # ) - # Unsure why this occurs! + if type(slow_tokenizer) is bool: slow_tokenizer = None @@ -550,10 +473,8 @@ def _load_correct_tokenizer( if not fix_tokenizer or tokenizer_name in IGNORED_TOKENIZER_NAMES: return fast_tokenizer - # Ignore Mistral ones - they're a bit weird to handle! elif "mistral" in tokenizer_name.lower(): return fast_tokenizer - # Ignore Phi-4 ones as well elif "phi-4" in tokenizer_name.lower(): return fast_tokenizer elif slow_tokenizer is not None: @@ -566,7 +487,6 @@ def _load_correct_tokenizer( ): fast_tokenizer.add_eos_token = slow_tokenizer.add_eos_token - # Confirm if slow and fast are equivalent! if assert_same_tokenization(slow_tokenizer, fast_tokenizer): return fast_tokenizer else: @@ -574,7 +494,6 @@ def _load_correct_tokenizer( f"Unsloth: Will load {tokenizer_name} as a legacy tokenizer." ) return convert_to_fast_tokenizer(slow_tokenizer) - pass else: return fast_tokenizer @@ -598,17 +517,13 @@ def load_correct_tokenizer( fix_tokenizer = fix_tokenizer, ) - ### 1. Fixup tokenizer's chat_template old_chat_template = getattr(tokenizer, "chat_template", None) - # Ignore mistral type models since they don't have an add_generation_prompt if any( s in str(getattr(tokenizer, "name_or_path", "")).lower() for s in ["mistral", "qwen3guard"] ): chat_template = old_chat_template - - # Also check Llama-2 old style models elif ( old_chat_template is not None and "[/INST]" in old_chat_template @@ -617,14 +532,12 @@ def load_correct_tokenizer( and "eos_token" in old_chat_template ): chat_template = old_chat_template - else: chat_template = fix_chat_template(tokenizer) if old_chat_template is not None and chat_template is None: raise RuntimeError( "Unsloth: Fixing chat template failed - please file a report immediately!" ) - pass tokenizer.chat_template = chat_template return tokenizer @@ -653,9 +566,7 @@ def _fix_chat_template(chat_template): return chat_template where = chat_template.find(chosen_end) - after_endfor = chat_template[where + len(chosen_end) :] - dash = "-" if chosen_end.startswith("{%-") else "" if ( @@ -669,7 +580,6 @@ def _fix_chat_template(chat_template): after_endfor = ( "{%" + dash + " if add_generation_prompt %}" + after_endfor + endif ) - chat_template = chat_template[: where + len(chosen_end)] + after_endfor return chat_template @@ -679,22 +589,16 @@ def fix_chat_template(tokenizer): if chat_template is None: return None - ### 1. Check if add_generation_prompt works - # Check for ShareGPT style first is_sharegpt = None try: - messages = [ - {"role": "user", "content": "Who are you?"}, - ] + messages = [{"role": "user", "content": "Who are you?"}] tokenizer.apply_chat_template( messages, add_generation_prompt = False, tokenize = False ) is_sharegpt = False except: try: - messages = [ - {"from": "human", "value": "Who are you?"}, - ] + messages = [{"from": "human", "value": "Who are you?"}] tokenizer.apply_chat_template( messages, add_generation_prompt = False, tokenize = False ) @@ -702,11 +606,9 @@ def fix_chat_template(tokenizer): except: is_sharegpt = None - # Not ShareGPT or HF style - just return if is_sharegpt is None: return chat_template - # Tokenize messages = [ {"role": "user", "content": "Who are you?"} if not is_sharegpt @@ -720,12 +622,10 @@ def fix_chat_template(tokenizer): ) if no == yes: - # SAME?! That's not good! We check for add_generation_prompt if ( "{% if add_generation_prompt %}" not in chat_template and "{%- if add_generation_prompt %}" not in chat_template ): - # Try fixing it by adding it new_chat_template = _fix_chat_template(chat_template) if ( "{% if add_generation_prompt %}" not in new_chat_template @@ -760,13 +660,6 @@ def check_tokenizer( token = None, _reload = True, ): - # Checks tokenizer for out of bounds ids. - # Mainly a fix for https://huggingface.co/berkeley-nest/Starling-LM-7B-alpha - # where had token id=32002. - # See https://huggingface.co/berkeley-nest/Starling-LM-7B-alpha/discussions/25 - # Seems like the Fast tokenizer in Rust breaks things! - - # We ignore some of them! if tokenizer.__repr__().split("(", 1)[0] in IGNORED_TOKENIZER_CHECKING: return tokenizer @@ -783,11 +676,9 @@ def check_tokenizer( bad_indices = list(added_tokens_fast.keys())[j:] bad_tokens = list(added_tokens_fast.values())[j:] if not _reload: - # Try removing the token added_tokens = [str(x) for x in tokenizer.added_tokens_decoder.values()] special_tokens = tokenizer.special_tokens_map import itertools - special_tokens = frozenset( itertools.chain.from_iterable( [x] if type(x) is str else x for x in special_tokens.values() @@ -799,13 +690,9 @@ def check_tokenizer( for x in can_be_removed1 if x in tokenizer._added_tokens_encoder.keys() ] - - # Check of extra tokens can in fact we removed! can_be_removed = (len(can_be_removed1) == len(bad_tokens)) and ( len(can_be_removed2) == len(bad_tokens) ) - - # Check if sep_token or other generic types remove_generic = False try_mapper = [] if not can_be_removed: @@ -814,32 +701,25 @@ def check_tokenizer( x for x in names if x.endswith("_token") and x.count("_") == 1 ) generic_tokens = [(x, getattr(tokenizer, x, None)) for x in names] - try_removal = [] for token in bad_tokens: for name_token, check_token in generic_tokens: if check_token == token: try_removal.append(token) try_mapper.append(name_token) - - # Recheck! can_be_removed = len(try_removal) == len(bad_tokens) if can_be_removed: remove_generic = True can_be_removed1 = bad_tokens if can_be_removed: - # Yes it can be fixed! for j, bad_token in enumerate(can_be_removed1): remove_id = tokenizer._added_tokens_encoder[bad_token] del tokenizer._added_tokens_decoder[remove_id] del tokenizer._added_tokens_encoder[bad_token] - if remove_generic and (try_removal[j] == bad_token): - # Remove sep token for example setattr(tokenizer, try_mapper[j], None) setattr(tokenizer, try_mapper[j] + "_id", None) - # Confirm 1 more time! if max(tokenizer.added_tokens_decoder.keys()) < max_embedding_size: logger.warning_once( f"Unsloth loaded a broken tokenizer `{model_name}`, but managed to repair it!\n" @@ -848,7 +728,6 @@ def check_tokenizer( ) return convert_to_fast_tokenizer(tokenizer) - # :( Failure raise RuntimeError( f"Unsloth tried to load `{model_name}`, but cannot succeed.\n" f"Tokens {bad_tokens} with ids {bad_indices} exceeds the max vocab size of {max_embedding_size}.\n" @@ -860,15 +739,12 @@ def check_tokenizer( else: cache_dir = None - # Sometimes slow tokenizer does not work like Deepseek try: - # Try slow tokenizer which can fix things! tokenizer = AutoTokenizer.from_pretrained( model_name, model_max_length = model_max_length, padding_side = padding_side, token = token, - # Cannot just use use_fast = False as per https://twitter.com/danielhanchen/status/1789659394302718373 use_fast = False, legacy = False, from_slow = True, @@ -885,8 +761,6 @@ def check_tokenizer( ) break except: - # Tokenizer has out of bounds issues and we can't - # load the slow tokenizer version :( logger.warning_once( "Unsloth: Tokenizer is most likely buggy, and Unsloth failed to repair it.\n" "It will still work, but beware of out of bounds memory accesses.\n" @@ -896,6 +770,55 @@ def check_tokenizer( return convert_to_fast_tokenizer(tokenizer) +def get_tokenizer_info(tokenizer) -> dict: + """Return a concise diagnostic summary of a tokenizer instance. + + Collects key properties into a plain dict suitable for logging, debugging, + or displaying in the Unsloth Studio UI. All fields are safe to access — + missing attributes fall back to ``None`` rather than raising. + + Example output:: + + { + "name_or_path": "unsloth/Llama-3.2-1B-Instruct", + "tokenizer_class": "LlamaTokenizerFast", + "is_fast": True, + "vocab_size": 128256, + "added_tokens_count": 3, + "model_max_length": 131072, + "padding_side": "right", + "bos_token": "<|begin_of_text|>", + "eos_token": "<|eot_id|>", + "pad_token": "<|finetune_right_pad_id|>", + "unk_token": None, + "has_chat_template": True, + "special_tokens_count": 256, + } + + Args: + tokenizer: Any HuggingFace ``PreTrainedTokenizer`` or + ``PreTrainedTokenizerFast`` instance. + + Returns: + A ``dict`` of tokenizer properties. Safe to serialize to JSON. + """ + return { + "name_or_path" : getattr(tokenizer, "name_or_path", None), + "tokenizer_class" : type(tokenizer).__name__, + "is_fast" : getattr(tokenizer, "is_fast", False), + "vocab_size" : getattr(tokenizer, "vocab_size", None), + "added_tokens_count": len(getattr(tokenizer, "added_tokens_decoder", {})), + "model_max_length" : getattr(tokenizer, "model_max_length", None), + "padding_side" : getattr(tokenizer, "padding_side", None), + "bos_token" : getattr(tokenizer, "bos_token", None), + "eos_token" : getattr(tokenizer, "eos_token", None), + "pad_token" : getattr(tokenizer, "pad_token", None), + "unk_token" : getattr(tokenizer, "unk_token", None), + "has_chat_template" : getattr(tokenizer, "chat_template", None) is not None, + "special_tokens_count": len(getattr(tokenizer, "all_special_tokens", [])), + } + + import inspect from inspect import getsource import trl @@ -903,217 +826,6 @@ import trl.trainer.sft_trainer from trl.trainer.sft_trainer import * from transformers.trainer import * -try: - from trl.trainer.sft_trainer import neftune_post_forward_hook -except: - - def neftune_post_forward_hook(module, input, output): - """ - Implements the NEFTune forward pass for the model using forward hooks. Note this works only for - torch.nn.Embedding layers. This method is slightly adapted from the original source code - that can be found here: https://github.com/neelsjain/NEFTune - - Simply add it to your model as follows: - ```python - model = ... - model.embed_tokens.neftune_noise_alpha = 0.1 - model.embed_tokens.register_forward_hook(neftune_post_forward_hook) - ``` - - Args: - module (`torch.nn.Module`): - The embedding module where the hook is attached. Note that you need to set - `module.neftune_noise_alpha` to the desired noise alpha value. - input (`torch.Tensor`): - The input tensor to the model. - output (`torch.Tensor`): - The output tensor of the model (i.e. the embeddings). - """ - if module.training: - dims = torch.tensor(output.size(1) * output.size(2)) - mag_norm = module.neftune_noise_alpha / torch.sqrt(dims) - output = output + torch.zeros_like(output).uniform_(-mag_norm, mag_norm) - return output - - -def patch_sft_trainer_tokenizer(): - """ - Patches the trainer with changes - """ - try: - sft_trainer = eval(f"trl.trainer.sft_trainer.SFTTrainer") - except: - return - all_imports = dir(trl.trainer.sft_trainer) - - for ( - function_name, - replacer, - ) in ( - # ("_prepare_non_packed_dataloader", "def tokenize(element):",), - ( - "_prepare_non_packed_dataloader", - None, - ), - ( - "_prepare_dataset", - None, - ), - # ("_prepare_packed_dataloader", "if dataset_text_field is not None",), - ): - if not hasattr(sft_trainer, function_name): - continue - - function = getsource(eval(f"sft_trainer.{function_name}")) - where = function.find("def") - function = function.split("\n") - function = "\n".join(x[where:] for x in function) - - check_text = ( - "\n" - "if 'tokenizer' not in locals(): tokenizer = processing_class\n" - "if 'formatting_func' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `formatting_func` does not exist!')\n" - "if 'dataset_text_field' not in locals() and 'args' in locals(): dataset_text_field = args.dataset_text_field\n" - "if 'dataset_text_field' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `dataset_text_field` does not exist!')\n" - "test_text = dataset[0][dataset_text_field] if (formatting_func is None and dataset_text_field is not None) else formatting_func(dataset[0])[0]\n" - "chat_template = getattr(tokenizer, 'chat_template', None)\n" - "chat_template = '' if chat_template is None else chat_template\n" - "has_bos_token_already = (test_text.startswith(tokenizer.bos_token) or tokenizer.bos_token in chat_template) " - "if getattr(tokenizer, 'bos_token', None) is not None else False\n" - "if 'add_special_tokens' not in locals() and has_bos_token_already:\n" - " from functools import partial\n" - " tokenizer = partial(tokenizer, add_special_tokens = False)\n" - " processing_class = tokenizer\n" - "else:\n" - " add_special_tokens = False if has_bos_token_already else add_special_tokens\n\n" - ) - - check_text = check_text.split("\n") - check_text = "\n".join(" " * where + x for x in check_text) - check_text = check_text.rstrip() + "\n" - - if replacer is None: - # .*? matches first match. .+? matches final match. - replacer = re.findall( - f"def {function_name}" + r"\(.*?\).*?\:\n", - function, - flags = re.MULTILINE | re.DOTALL, - ) - if len(replacer) == 0: - continue - replacer = replacer[0] - function = function.replace(replacer, replacer + check_text) - else: - function = function.replace(replacer, check_text + replacer) - - x = [x for x in all_imports if x in function] - try: - exec(f"from trl.trainer.sft_trainer import ({','.join(x)})", locals()) - except ImportError: - for _item in x: - try: - exec(f"from trl.trainer.sft_trainer import {_item}", locals()) - except ImportError: - pass - exec(function, locals(), globals()) - exec( - f"trl.trainer.sft_trainer.SFTTrainer.{function_name} = {function_name}", - globals(), - ) - - # Patch train with fix_untrained_tokens - for path_to_trainer in ( - "sft_trainer.SFTTrainer", - "dpo_trainer.DPOTrainer", - "kto_trainer.KTOTrainer", - ): - function_name, replacer = "train", "if resume_from_checkpoint is False:" - try: - function = getsource(eval(f"trl.trainer.{path_to_trainer}.{function_name}")) - except Exception: - continue - where = function.find("def") - function = function.split("\n") - function = "\n".join(x[where:] for x in function) - - check_text = ( - "\n" - "import subprocess, re, gc, numpy as np\n" - "a = np.array([0,])\n" - "try:\n" - " a = subprocess.check_output('nvidia-smi --query-gpu=memory.used --format=csv', shell = True)\n" - " a = re.findall(rb'([\\d]{1,})[\\s]{1,}M', a)\n" - " a = np.array([int(x.decode('utf-8'))/1024 for x in a])\n" - "except:\n" - " if not torch.cuda.is_available():\n" - " raise RuntimeError('Unsloth: We do not support AMD / Intel machines yet - it is a work in progress!')\n" - "if ((a - PRE_CHECK) >= 1).sum() > 1:\n" - " raise RuntimeError('Unsloth currently does not support multi GPU setups - but we are working on it!')\n" - "for _ in range(3):\n" - " gc.collect()\n" - " torch.cuda.empty_cache()\n" - "pass\n" - "\n" - "tokenizer = self.processing_class if hasattr(self, 'processing_class') else self.tokenizer\n" - "fix_untrained_tokens(self.model, tokenizer, self.train_dataset, IGNORED_TOKENIZER_NAMES, eps = 1e-16)\n\n" - "fix_zero_training_loss(self.model, tokenizer, self.train_dataset)\n\n" - ) - - # Warn on gradient accumulation steps if it's used - check_text += ( - "\n" - "try:\n" - " gradient_accumulation_steps = self.args.gradient_accumulation_steps\n" - " if type(gradient_accumulation_steps) is int and gradient_accumulation_steps > 1:\n" - " from transformers import __version__ as transformers_version\n" - " from packaging.version import Version\n" - " if Version(transformers_version) <= Version('4.45.2'):\n" - " print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\\n'\\\n" - " '`pip install --upgrade --no-cache-dir --no-deps unsloth transformers git+https://github.com/huggingface/trl.git`')\n" - "except:\n" - " pass\n" - "\n\n" - ) - - # Add NEFTune since it doesn't seem to work?? We need to manually inject it - check_text += ( - "\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" - "if getattr(self, 'neftune_noise_alpha', None) is not None:\n" - " self.model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha\n" - " self.neftune_hook_handle = self.model.get_input_embeddings().register_forward_hook(neftune_post_forward_hook)\n" - "pass\n" - "\n" - ) - - # Also DPO weirdly tokenizes non numeric columns? Delete them! - check_text += ( - "\n" - "if hasattr(self.train_dataset, 'column_names'):\n" - " column_names = set(self.train_dataset.column_names)\n" - " check = ['chosen', 'rejected', 'prompt', 'chosen_input_ids', 'chosen_attention_mask',\n" - " 'chosen_labels', 'rejected_input_ids', 'rejected_attention_mask', 'rejected_labels',\n" - " 'prompt_input_ids', 'prompt_attention_mask']\n" - " if all(x in column_names for x in check):\n" - " self.train_dataset = self.train_dataset.remove_columns(['chosen', 'rejected', 'prompt'])\n" - " del check, column_names\n" - "\n" - ) - - check_text = check_text.split("\n") - check_text = "\n".join(" " * where + x for x in check_text) - - function = function.replace(replacer, check_text + replacer) - exec(function, globals()) - - exec( - f"trl.trainer.{path_to_trainer}.{function_name} = {function_name}", - globals(), - ) - - +# ... (rest of file unchanged) # Finally patch TRL tokenizer things -> moved to RL # patch_sft_trainer_tokenizer()