Tokenizer overwritten

This commit is contained in:
Daniel Han-Chen 2024-03-10 04:31:40 +11:00
commit f887080a6d
3 changed files with 23 additions and 12 deletions

View file

@ -916,6 +916,7 @@ class FastLlamaModel:
rope_scaling = None, rope_scaling = None,
fix_tokenizer = True, fix_tokenizer = True,
model_patcher = None, model_patcher = None,
tokenizer_name = None,
**kwargs, **kwargs,
): ):
if model_patcher is None: model_patcher = FastLlamaModel if model_patcher is None: model_patcher = FastLlamaModel
@ -978,18 +979,17 @@ class FastLlamaModel:
max_position_embeddings = max_position_embeddings, max_position_embeddings = max_position_embeddings,
**kwargs, **kwargs,
) )
# Counteract saved tokenizers
tokenizer_name = model_name if tokenizer_name is None else tokenizer_name
tokenizer = AutoTokenizer.from_pretrained( tokenizer = AutoTokenizer.from_pretrained(
model_name, tokenizer_name,
model_max_length = max_position_embeddings, model_max_length = max_position_embeddings,
padding_side = "right", padding_side = "right",
token = token, token = token,
) )
print(tokenizer)
print(tokenizer.chat_template)
model, tokenizer = patch_tokenizer(model, tokenizer) model, tokenizer = patch_tokenizer(model, tokenizer)
print(tokenizer)
print(tokenizer.chat_template)
model = model_patcher.post_patch(model) model = model_patcher.post_patch(model)
# Patch up QKV / O and MLP # Patch up QKV / O and MLP

View file

@ -18,7 +18,7 @@ from transformers import AutoConfig
from transformers import __version__ as transformers_version from transformers import __version__ as transformers_version
from peft import PeftConfig, PeftModel 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
import os
# https://github.com/huggingface/transformers/pull/26037 allows 4 bit loading! # https://github.com/huggingface/transformers/pull/26037 allows 4 bit loading!
major, minor = transformers_version.split(".")[:2] major, minor = transformers_version.split(".")[:2]
@ -79,14 +79,12 @@ class FastLanguageModel(FastLlamaModel):
): ):
old_model_name = model_name old_model_name = model_name
model_name = _get_model_name(model_name, load_in_4bit) model_name = _get_model_name(model_name, load_in_4bit)
print(model_name)
# First check if it's a normal model via AutoConfig # First check if it's a normal model via AutoConfig
is_peft = False is_peft = False
try: try:
model_config = AutoConfig.from_pretrained(model_name, token = token) model_config = AutoConfig.from_pretrained(model_name, token = token)
is_peft = False is_peft = False
print(model_config)
except: except:
try: try:
# Most likely a PEFT model # Most likely a PEFT model
@ -98,7 +96,6 @@ class FastLanguageModel(FastLlamaModel):
model_name = _get_model_name(peft_config.base_model_name_or_path, load_in_4bit) model_name = _get_model_name(peft_config.base_model_name_or_path, load_in_4bit)
model_config = AutoConfig.from_pretrained(model_name, token = token) model_config = AutoConfig.from_pretrained(model_name, token = token)
is_peft = True is_peft = True
print(model_config)
pass pass
model_type = model_config.model_type model_type = model_config.model_type
@ -121,7 +118,16 @@ class FastLanguageModel(FastLlamaModel):
) )
pass pass
print(model_name) # Check if this is local model since the tokenizer gets overwritten
if os.path.exists(os.path.join(old_model_name, "tokenizer_config.json")) and \
os.path.exists(os.path.join(old_model_name, "tokenizer.json")) and \
os.path.exists(os.path.join(old_model_name, "special_tokens_map.json")):
tokenizer_name = old_model_name
else:
tokenizer_name = None
pass
model, tokenizer = dispatch_model.from_pretrained( model, tokenizer = dispatch_model.from_pretrained(
model_name = model_name, model_name = model_name,
max_seq_length = max_seq_length, max_seq_length = max_seq_length,
@ -132,6 +138,7 @@ class FastLanguageModel(FastLlamaModel):
rope_scaling = rope_scaling, rope_scaling = rope_scaling,
fix_tokenizer = fix_tokenizer, fix_tokenizer = fix_tokenizer,
model_patcher = dispatch_model, model_patcher = dispatch_model,
tokenizer_name = tokenizer_name,
*args, **kwargs, *args, **kwargs,
) )

View file

@ -294,6 +294,7 @@ class FastMistralModel(FastLlamaModel):
rope_scaling = None, # Mistral does not support RoPE scaling rope_scaling = None, # Mistral does not support RoPE scaling
fix_tokenizer = True, fix_tokenizer = True,
model_patcher = None, model_patcher = None,
tokenizer_name = None,
**kwargs, **kwargs,
): ):
if model_patcher is None: model_patcher = FastMistralModel if model_patcher is None: model_patcher = FastMistralModel
@ -354,8 +355,11 @@ class FastMistralModel(FastLlamaModel):
# rope_scaling = rope_scaling, # rope_scaling = rope_scaling,
**kwargs, **kwargs,
) )
# Counteract saved tokenizers
tokenizer_name = model_name if tokenizer_name is None else tokenizer_name
tokenizer = AutoTokenizer.from_pretrained( tokenizer = AutoTokenizer.from_pretrained(
model_name, tokenizer_name,
model_max_length = max_position_embeddings, model_max_length = max_position_embeddings,
padding_side = "right", padding_side = "right",
token = token, token = token,