From 5393e9e00a1e2019144698d90035ae21e03325c7 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 15 Aug 2024 00:31:30 -0700 Subject: [PATCH] Bug Fixes (#920) * Update pyproject.toml * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update _utils.py * Update _utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * fix_tokenizer * Update tokenizer_utils.py * Update tokenizer_utils.py * Update save.py * Update save.py * Update save.py * Update save.py * Update save.py * Update loader.py * Update pyproject.toml * Update _utils.py * Update gemma2.py * Update gemma2.py * Update _utils.py * gemma 2 mask * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update _utils.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update _utils.py * Update llama.py * Update llama.py * Update llama.py * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Torch 2.4 Xformers 0.0.27post2 * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Gemma 2 fixes * Update gemma2.py * Update llama.py * Update llama.py * Update save.py * Update save.py * Update llama.py * Update cross_entropy_loss.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Update dpo.py * Providing more flexibility for users to customize their llama when using LoRA (#910) * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update chat_templates.py * return model * Update tokenizer_utils.py * Update chat_templates.py * Update tokenizer_utils.py * Train on completions * load_in_4bit=False broken --------- Co-authored-by: Po-Lung Wang --- unsloth/models/llama.py | 6 +++++- unsloth/models/loader.py | 19 +++++++++++-------- 2 files changed, 16 insertions(+), 9 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 6a111c9344..6139115f67 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -1387,11 +1387,15 @@ class FastLlamaModel: # RoPE Scaling's max_position_embeddings must be updated max_position_embeddings = max(max_seq_length, model_max_seq_length) kwargs.pop("attn_implementation", None); # No need since we auto call it + + # Cannot be None, since HF now checks for the config + if load_in_4bit: kwargs["quantization_config"] = bnb_config + model = AutoModelForCausalLM.from_pretrained( model_name, device_map = device_map, torch_dtype = dtype, - quantization_config = bnb_config, + # quantization_config = bnb_config, token = token, max_position_embeddings = max_position_embeddings, trust_remote_code = trust_remote_code, diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index cce22aebf6..ad1098edac 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -42,10 +42,11 @@ def __get_model_name( INT_TO_FLOAT_MAPPER = None, FLOAT_TO_INT_MAPPER = None, ): - model_name = str(model_name) - if not SUPPORTS_FOURBIT and model_name.lower() in INT_TO_FLOAT_MAPPER: - model_name = INT_TO_FLOAT_MAPPER[model_name.lower()] + lower_model_name = model_name.lower() + + if not SUPPORTS_FOURBIT and lower_model_name in INT_TO_FLOAT_MAPPER: + model_name = INT_TO_FLOAT_MAPPER[lower_model_name] logger.warning_once( f"Unsloth: Your transformers version of {transformers_version} does not support native "\ f"4bit loading.\nThe minimum required version is 4.37.\n"\ @@ -55,16 +56,18 @@ def __get_model_name( ) return model_name - elif not load_in_4bit and model_name.lower() in INT_TO_FLOAT_MAPPER: - new_model_name = INT_TO_FLOAT_MAPPER[model_name.lower()] + elif not load_in_4bit and lower_model_name in INT_TO_FLOAT_MAPPER: + new_model_name = INT_TO_FLOAT_MAPPER[lower_model_name] # logger.warning_once( # f"Unsloth: You passed in `{model_name}` which is a 4bit model, yet you set\n"\ # f"`load_in_4bit = False`. We shall load `{new_model_name}` instead." # ) return new_model_name - - elif load_in_4bit and SUPPORTS_FOURBIT and model_name.lower() in FLOAT_TO_INT_MAPPER: - new_model_name = FLOAT_TO_INT_MAPPER[model_name.lower()] + elif not load_in_4bit and lower_model_name in FLOAT_TO_INT_MAPPER: + new_model_name = FLOAT_TO_INT_MAPPER[lower_model_name] + return new_model_name + elif load_in_4bit and SUPPORTS_FOURBIT and lower_model_name in FLOAT_TO_INT_MAPPER: + new_model_name = FLOAT_TO_INT_MAPPER[lower_model_name] # logger.warning_once( # f"Unsloth: You passed in `{model_name}` and `load_in_4bit = True`.\n"\ # f"We shall load `{new_model_name}` for 4x faster loading."