From 36473e2d6ef8d11fe7cf5fc798c41a3e8403dc64 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Mon, 18 Mar 2024 04:18:15 +1100 Subject: [PATCH] Fix lm_head, embed_tokens (#258) * Update gemma.py * Update gemma.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * rope * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * llama * Update llama.py * gemma * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update save.py * RoPE * Update llama.py * Update llama.py * Update llama.py * Update gemma.py * correct_dtype * Update gemma.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Chat Templates * Update README.md * Update README.md * Update llama.py * DoRA * Update _utils.py * Update chat_templates.py * Update llama.py * Hotfix - fix DoRA, Gemma prompt template (#202) (#203) * Update save.py * saving * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update __init__.py * Update save.py * Update save.py * Update save.py * save * trainer * spaces * original * Gemma * Update pyproject.toml * Update mapper.py * Update fast_lora.py * FastGemmaModel * model_type * Update llama.py * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update llama.py * Update fast_lora.py * Update llama.py * Update llama.py * Update cross_entropy_loss.py * Update llama.py * Update llama.py * gemma * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update fast_lora.py * Update fast_lora.py * Fast CE Loss * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * CE * Update llama.py * Update llama.py * Update cross_entropy_loss.py * Update geglu.py * Update cross_entropy_loss.py * revert * Update llama.py * Update llama.py * norm * Update gemma.py * Update gemma.py * position_ids * Update gemma.py * Update gemma.py * pos * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * revert * revert * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * rope * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * llama * Update llama.py * gemma * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update save.py * RoPE * Update llama.py * Update llama.py * Update llama.py * Update gemma.py * correct_dtype * Update gemma.py * Update cross_entropy_loss.py * Update cross_entropy_loss.py * Chat Templates * Update README.md * Update README.md * Update llama.py * DoRA * Update _utils.py * Update chat_templates.py * Update pyproject.toml * Small fixes * Update pyproject.toml * Approx gelu * Update geglu.py * Approx gelu * Update llama.py * Update __init__.py * Update __init__.py * Update _utils.py * Update geglu.py * Update gemma.py * Update rms_layernorm.py * Update rms_layernorm.py * Update rms_layernorm.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * Fix Gemma merging * Update rms_layernorm.py * Update gemma.py * Update pyproject.toml * Layernorms * Gemma precision * Update gemma.py * sqrt * Update gemma.py * Update save.py * RoPE and Gemma precision * Update rms_layernorm.py * Fix warning * Update chat_templates.py * Update chat_templates.py * Update save.py * Update save.py * Update save.py * Update chat_templates.py * Update llama.py * model_name * Update loader.py * Tokenizer overwritten * Update llama.py * Update llama.py * Update llama.py * Update save.py * Accuracy * Revert * Update save.py * Update fast_lora.py * Update fast_lora.py * Update fast_lora.py * Update fast_lora.py * Update fast_lora.py * Update chat_templates.py * Update save.py * Update save.py * Update llama.py * Update llama.py * Account for DoRA * Update llama.py * Update save.py * GGUF incorrect * Update save.py * Update pyproject.toml * kaggle new * Update pyproject.toml * Update pyproject.toml * upcasting * Fix Colab * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update chat_templates.py * Update chat_templates.py * Update chat_templates.py * Update chat_templates.py * Update chat_templates.py * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update rope_embedding.py * Update rope_embedding.py * Fix bugs * Update fast_lora.py * Update fast_lora.py * Update README.md * Update README.md * GGUF * Update save.py * Update save.py * Update save.py * Update save.py * Update README.md * Update README.md * Bugs * Update fast_lora.py * Update pyproject.toml * Update fast_lora.py * Update __init__.py * Update fast_lora.py * dtype * Update llama.py * Update llama.py * Update llama.py * dtype * Update mistral.py * trust_remote_code --- unsloth/models/llama.py | 27 ++++++++++++++++++--------- unsloth/models/loader.py | 2 ++ unsloth/models/mistral.py | 10 +++++++--- 3 files changed, 27 insertions(+), 12 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index d83d9b76f2..8b86feb35e 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -505,11 +505,10 @@ def LlamaModel_fast_forward( position_ids = position_ids.repeat((batch_size, 1)) pass - # embed positions + # Embed positions if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) - # Downcast to the correct dtype ie float32 to float16 inputs_embeds = inputs_embeds.to(self.config.torch_dtype) # Normalized from Gemma @@ -759,6 +758,7 @@ def CausalLM_fast_forward(fast_forward_inference): else: logits = self.lm_head(hidden_states) pass + logits = logits.to(self.config.torch_dtype) loss = None if labels is not None: @@ -929,6 +929,7 @@ class FastLlamaModel: fix_tokenizer = True, model_patcher = None, tokenizer_name = None, + trust_remote_code = False, **kwargs, ): if model_patcher is None: model_patcher = FastLlamaModel @@ -989,6 +990,7 @@ class FastLlamaModel: token = token, rope_scaling = rope_scaling, max_position_embeddings = max_position_embeddings, + trust_remote_code = trust_remote_code, **kwargs, ) @@ -996,9 +998,10 @@ class FastLlamaModel: tokenizer_name = model_name if tokenizer_name is None else tokenizer_name tokenizer = AutoTokenizer.from_pretrained( tokenizer_name, - model_max_length = max_position_embeddings, - padding_side = "right", - token = token, + model_max_length = max_position_embeddings, + padding_side = "right", + token = token, + trust_remote_code = trust_remote_code, ) model, tokenizer = patch_tokenizer(model, tokenizer) @@ -1338,7 +1341,6 @@ class FastLlamaModel: "We shall do it for you!" ) train_lm_head = True - model.model.embed_tokens.to(torch.float32, non_blocking = True) elif module == "embed_tokens": logger.warning_once( @@ -1346,7 +1348,6 @@ class FastLlamaModel: "We shall do it for you!" ) train_embed_tokens = True - model.lm_head.to(torch.float32, non_blocking = True) else: assert(module in accepted_modules) @@ -1388,9 +1389,17 @@ class FastLlamaModel: # Now patch lm_head and embed_tokens if train_embed_tokens: - model.model.model.embed_tokens.requires_grad_(True) + print("Unsloth: Casting embed_tokens to float32") + assert(hasattr(model.model.model.embed_tokens, "modules_to_save")) + model.model.model.embed_tokens.modules_to_save.default.to(torch.float32) + model.model.model.embed_tokens.modules_to_save.default.requires_grad_(True) + pass + if train_lm_head: - model.model.lm_head.requires_grad_(True) + print("Unsloth: Casting lm_head to float32") + assert(hasattr(model.model.lm_head, "modules_to_save")) + model.model.lm_head.modules_to_save.default.to(torch.float32) + model.model.lm_head.modules_to_save.default.requires_grad_(True) pass return model diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index 47b568ae2a..29d25f3c20 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -74,6 +74,7 @@ class FastLanguageModel(FastLlamaModel): device_map = "sequential", rope_scaling = None, fix_tokenizer = True, + trust_remote_code = False, use_gradient_checkpointing = True, *args, **kwargs, ): @@ -139,6 +140,7 @@ class FastLanguageModel(FastLlamaModel): fix_tokenizer = fix_tokenizer, model_patcher = dispatch_model, tokenizer_name = tokenizer_name, + trust_remote_code = trust_remote_code, *args, **kwargs, ) diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index c1e39e4a2e..2db71f0288 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -230,6 +230,7 @@ def MistralForCausalLM_fast_forward( else: logits = self.lm_head(hidden_states) pass + logits = logits.to(self.config.torch_dtype) loss = None if labels is not None: @@ -295,6 +296,7 @@ class FastMistralModel(FastLlamaModel): fix_tokenizer = True, model_patcher = None, tokenizer_name = None, + trust_remote_code = False, **kwargs, ): if model_patcher is None: model_patcher = FastMistralModel @@ -353,6 +355,7 @@ class FastMistralModel(FastLlamaModel): quantization_config = bnb_config, token = token, # rope_scaling = rope_scaling, + trust_remote_code = trust_remote_code, **kwargs, ) @@ -360,9 +363,10 @@ class FastMistralModel(FastLlamaModel): tokenizer_name = model_name if tokenizer_name is None else tokenizer_name tokenizer = AutoTokenizer.from_pretrained( tokenizer_name, - model_max_length = max_position_embeddings, - padding_side = "right", - token = token, + model_max_length = max_position_embeddings, + padding_side = "right", + token = token, + trust_remote_code = trust_remote_code, ) model, tokenizer = patch_tokenizer(model, tokenizer)