From c90bb8dc326381bde6dcbeda75fd55d65a390b61 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Thu, 15 Aug 2024 01:15:35 -0700 Subject: [PATCH] Fix mapping (#921) * 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 * Update llama.py * MAP_TO_UNSLOTH_16bit * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update mapper.py * Update mapper.py * works! --------- Co-authored-by: Po-Lung Wang --- unsloth/models/llama.py | 2 +- unsloth/models/loader.py | 39 +++++++++++++++++++++++++-------------- unsloth/models/mapper.py | 13 +++++++++++-- 3 files changed, 37 insertions(+), 17 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 6139115f67..6a23335c8c 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -1390,7 +1390,7 @@ class FastLlamaModel: # 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, diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index ad1098edac..e260017fb9 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -19,7 +19,7 @@ from .qwen2 import FastQwen2Model from transformers import AutoConfig from transformers import __version__ as transformers_version from peft import PeftConfig, PeftModel -from .mapper import INT_TO_FLOAT_MAPPER, FLOAT_TO_INT_MAPPER +from .mapper import INT_TO_FLOAT_MAPPER, FLOAT_TO_INT_MAPPER, MAP_TO_UNSLOTH_16bit import os # https://github.com/huggingface/transformers/pull/26037 allows 4 bit loading! @@ -39,13 +39,15 @@ pass def __get_model_name( model_name, load_in_4bit = True, - INT_TO_FLOAT_MAPPER = None, - FLOAT_TO_INT_MAPPER = None, + INT_TO_FLOAT_MAPPER = None, + FLOAT_TO_INT_MAPPER = None, + MAP_TO_UNSLOTH_16bit = None, ): model_name = str(model_name) 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 "\ @@ -57,16 +59,21 @@ def __get_model_name( return model_name 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 not load_in_4bit and lower_model_name in FLOAT_TO_INT_MAPPER: - new_model_name = FLOAT_TO_INT_MAPPER[lower_model_name] + + elif not load_in_4bit and lower_model_name in MAP_TO_UNSLOTH_16bit: + + new_model_name = MAP_TO_UNSLOTH_16bit[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"\ @@ -86,12 +93,14 @@ def _get_new_mapper(): with requests.get(new_mapper, timeout = 3) as new_mapper: new_mapper = new_mapper.text new_mapper = new_mapper[new_mapper.find("__INT_TO_FLOAT_MAPPER"):] new_mapper = new_mapper\ - .replace("INT_TO_FLOAT_MAPPER", "NEW_INT_TO_FLOAT_MAPPER")\ - .replace("FLOAT_TO_INT_MAPPER", "NEW_FLOAT_TO_INT_MAPPER") + .replace("INT_TO_FLOAT_MAPPER", "NEW_INT_TO_FLOAT_MAPPER")\ + .replace("FLOAT_TO_INT_MAPPER", "NEW_FLOAT_TO_INT_MAPPER")\ + .replace("MAP_TO_UNSLOTH_16bit", "NEW_MAP_TO_UNSLOTH_16bit") + exec(new_mapper, globals()) - return NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER + return NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER, NEW_MAP_TO_UNSLOTH_16bit except: - return {}, {} + return {}, {}, {} pass pass @@ -100,17 +109,19 @@ def get_model_name(model_name, load_in_4bit = True): new_model_name = __get_model_name( model_name = model_name, load_in_4bit = load_in_4bit, - INT_TO_FLOAT_MAPPER = INT_TO_FLOAT_MAPPER, - FLOAT_TO_INT_MAPPER = FLOAT_TO_INT_MAPPER, + INT_TO_FLOAT_MAPPER = INT_TO_FLOAT_MAPPER, + FLOAT_TO_INT_MAPPER = FLOAT_TO_INT_MAPPER, + MAP_TO_UNSLOTH_16bit = MAP_TO_UNSLOTH_16bit, ) if new_model_name is None and model_name.count("/") == 1 and model_name[0].isalnum(): # Try checking if a new Unsloth version allows it! - NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER = _get_new_mapper() + NEW_INT_TO_FLOAT_MAPPER, NEW_FLOAT_TO_INT_MAPPER, NEW_MAP_TO_UNSLOTH_16bit = _get_new_mapper() upgraded_model_name = __get_model_name( model_name = model_name, load_in_4bit = load_in_4bit, - INT_TO_FLOAT_MAPPER = NEW_INT_TO_FLOAT_MAPPER, - FLOAT_TO_INT_MAPPER = NEW_FLOAT_TO_INT_MAPPER, + INT_TO_FLOAT_MAPPER = NEW_INT_TO_FLOAT_MAPPER, + FLOAT_TO_INT_MAPPER = NEW_FLOAT_TO_INT_MAPPER, + MAP_TO_UNSLOTH_16bit = NEW_MAP_TO_UNSLOTH_16bit, ) if upgraded_model_name is not None: raise NotImplementedError( diff --git a/unsloth/models/mapper.py b/unsloth/models/mapper.py index 57ba676585..b8259a073c 100644 --- a/unsloth/models/mapper.py +++ b/unsloth/models/mapper.py @@ -251,8 +251,9 @@ __INT_TO_FLOAT_MAPPER = \ ), } -INT_TO_FLOAT_MAPPER = {} -FLOAT_TO_INT_MAPPER = {} +INT_TO_FLOAT_MAPPER = {} +FLOAT_TO_INT_MAPPER = {} +MAP_TO_UNSLOTH_16bit = {} for key, values in __INT_TO_FLOAT_MAPPER.items(): INT_TO_FLOAT_MAPPER[key] = values[0] @@ -261,6 +262,14 @@ for key, values in __INT_TO_FLOAT_MAPPER.items(): FLOAT_TO_INT_MAPPER[value] = key pass + # Map to Unsloth version for 16bit versions + if len(values) == 2: + if values[0].startswith("unsloth"): + MAP_TO_UNSLOTH_16bit[values[1]] = values[0] + MAP_TO_UNSLOTH_16bit[values[1].lower()] = values[0] + pass + pass + # Get lowercased lowered_key = key.lower() INT_TO_FLOAT_MAPPER[lowered_key] = values[0].lower()