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 <Brownwang0426@gmail.com>
This commit is contained in:
parent
2a692ebb31
commit
c69fb285df
3 changed files with 37 additions and 17 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue