From 1c044da660810c422b32041cbfbd1519ff2db6e9 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 27 Oct 2024 19:06:57 -0700 Subject: [PATCH] Fix pad token --- unsloth/models/_utils.py | 27 ++++++++++++++-- unsloth/models/gemma.py | 7 ++--- unsloth/models/gemma2.py | 7 ++--- unsloth/models/llama.py | 68 ++++++++++++++++++++++++++++++---------- 4 files changed, 80 insertions(+), 29 deletions(-) diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index d92a938eee..46a5f45cfa 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -525,7 +525,7 @@ def patch_tokenizer(model, tokenizer): pass logger.warning_once( - f"Unsloth: unk_token = {old_unk_token} is the same as the EOS or BOS tokens.\n"\ + f"Unsloth: unk_token = {old_unk_token} is the same as the EOS or BOS tokens. "\ f"We fixed it by changing it to {tokenizer.unk_token}." ) pass @@ -610,13 +610,34 @@ def patch_tokenizer(model, tokenizer): tokenizer.add_special_tokens({"pad_token" : possible_pad_token}) tokenizer.pad_token = possible_pad_token if model is not None: - model.config.update({"pad_token_id" : tokenizer.pad_token_id}) + + # Edit all config with new pad token + current_model = model + while hasattr(model, "model") and hasattr(model, "config"): + current_model.config.update({"pad_token_id" : tokenizer.pad_token_id}) + current_model = current_model.model + if hasattr(model, "model") and hasattr(model, "config"): + current_model.config.update({"pad_token_id" : tokenizer.pad_token_id}) + pass + + # Generation edit pad token if getattr(model, "generation_config") is not None: model.generation_config.update(pad_token_id = tokenizer.pad_token_id) else: if model is not None: + if model.config.pad_token_id is None: - model.config.update({"pad_token_id" : tokenizer.pad_token_id}) + + # Edit all config with new pad token + current_model = model + while hasattr(model, "model") and hasattr(model, "config"): + current_model.config.update({"pad_token_id" : tokenizer.pad_token_id}) + current_model = model + if hasattr(model, "model") and hasattr(model, "config"): + current_model.config.update({"pad_token_id" : tokenizer.pad_token_id}) + pass + + # Generation edit pad token if getattr(model, "generation_config") is not None: model.generation_config.update(pad_token_id = tokenizer.pad_token_id) pass diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index 45f14c1131..1ec116b2ea 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -339,10 +339,7 @@ class FastGemmaModel(FastLlamaModel): @staticmethod - def post_patch(model): - # Patch model for Gemma - layers = model.model.layers - + def post_patch(model, tokenizer): # Torch.compile fails on embedding matrix?? # Workaround randomnly fixes it for torch versions < 2.2 model.model.embed_tokens = torch.nn.Embedding.from_pretrained(model.model.embed_tokens.weight) @@ -425,6 +422,6 @@ class FastGemmaModel(FastLlamaModel): for _ in range(3): gc.collect() torch.cuda.empty_cache() - return model + return model, tokenizer pass pass diff --git a/unsloth/models/gemma2.py b/unsloth/models/gemma2.py index bf40ea8a27..54d8f628cb 100644 --- a/unsloth/models/gemma2.py +++ b/unsloth/models/gemma2.py @@ -490,10 +490,7 @@ class FastGemma2Model(FastLlamaModel): @staticmethod - def post_patch(model): - # Patch model for Gemma - layers = model.model.layers - + def post_patch(model, tokenizer): # Torch.compile fails on embedding matrix?? # Workaround randomnly fixes it for torch versions < 2.2 model.model.embed_tokens = torch.nn.Embedding.from_pretrained(model.model.embed_tokens.weight) @@ -576,6 +573,6 @@ class FastGemma2Model(FastLlamaModel): for _ in range(3): gc.collect() torch.cuda.empty_cache() - return model + return model, tokenizer pass pass diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 9c9ea53752..044ea6e244 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -1621,7 +1621,7 @@ class FastLlamaModel: ) model, tokenizer = patch_tokenizer(model, tokenizer) - model = model_patcher.post_patch(model) + model, tokenizer = model_patcher.post_patch(model, tokenizer) # Patch up QKV / O and MLP for idx, layer in enumerate(model.model.layers): @@ -1827,27 +1827,63 @@ class FastLlamaModel: @staticmethod - def post_patch(model): - # Patch model - layers = model.model.layers - + def post_patch(model, tokenizer): # Torch.compile fails on embedding matrix?? - # Workaround randomnly fixes it for torch versions < 2. - model.set_input_embeddings(torch.nn.Embedding.from_pretrained(model.get_input_embeddings().weight)) + try: old_input_embedding = model.get_input_embeddings ().weight + except: return model, tokenizer + + # Maybe not all models have a lm_head? + try: old_output_embedding = model.get_output_embeddings().weight + except: old_output_embedding = torch.zeros(0) + + # Check for tied weights as well + is_tied = old_input_embedding.data_ptr() == old_output_embedding.data_ptr() + + # Check pad token's id -> we need to expand the embedding + if len(tokenizer) > old_input_embedding.shape[0]: + # Workaround randomnly fixes it for torch versions < 2. + requires_grad = old_input_embedding.requires_grad + old_input_embedding.requires_grad_(False) + old_input_embedding.resize_(len(tokenizer), old_input_embedding.shape[1]) + old_input_embedding.requires_grad_(requires_grad) + + # Fix up all vocab sizes + current_model = model + while hasattr(model, "model") and hasattr(model, "config"): + if hasattr(model.config, "vocab_size"): + current_model.config.update({"vocab_size" : len(tokenizer)}) + current_model = current_model.model + if hasattr(model, "model") and hasattr(model, "config"): + if hasattr(model.config, "vocab_size"): + current_model.config.update({"vocab_size" : len(tokenizer)}) + pass + pass + + model.set_input_embeddings( + torch.nn.Embedding.from_pretrained( + old_input_embedding, + padding_idx = getattr(model.config, "pad_token_id", None), + ) + ) model.config.update({"unsloth_version" : __version__}) # We also do this for the lm_head - lm_head = torch.nn.Linear(1, 1, bias = None) - del lm_head.weight - lm_head.weight = model.get_output_embeddings().weight - lm_head.in_features = lm_head.weight.shape[1] - lm_head.out_features = lm_head.weight.shape[0] - model.lm_head = lm_head + if old_output_embedding.numel() != 0: + requires_grad = old_output_embedding.requires_grad + lm_head = torch.nn.Linear(1, 1, bias = None) + del lm_head.weight + lm_head.weight = old_output_embedding if not is_tied else old_input_embedding + lm_head.in_features = lm_head.weight.shape[1] + lm_head.out_features = lm_head.weight.shape[0] + lm_head.weight.requires_grad_(requires_grad) + model.lm_head = lm_head + correct_dtype = lm_head.weight.dtype + else: + correct_dtype = old_input_embedding.dtype + pass # Also patch all dtypes - BnB seems to not allocate the correct type? # BnB default dtype seems to be float16! - correct_dtype = lm_head.weight.dtype - for name, module in model.named_modules(): if isinstance(module, (Bnb_Linear4bit, Peft_Linear4bit)): weight = module.weight @@ -1883,7 +1919,7 @@ class FastLlamaModel: for _ in range(3): gc.collect() torch.cuda.empty_cache() - return model + return model, tokenizer pass