Tokenizer overwritten
This commit is contained in:
parent
1c1461ae09
commit
f887080a6d
3 changed files with 23 additions and 12 deletions
|
|
@ -916,6 +916,7 @@ class FastLlamaModel:
|
|||
rope_scaling = None,
|
||||
fix_tokenizer = True,
|
||||
model_patcher = None,
|
||||
tokenizer_name = None,
|
||||
**kwargs,
|
||||
):
|
||||
if model_patcher is None: model_patcher = FastLlamaModel
|
||||
|
|
@ -978,18 +979,17 @@ class FastLlamaModel:
|
|||
max_position_embeddings = max_position_embeddings,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Counteract saved tokenizers
|
||||
tokenizer_name = model_name if tokenizer_name is None else tokenizer_name
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_name,
|
||||
tokenizer_name,
|
||||
model_max_length = max_position_embeddings,
|
||||
padding_side = "right",
|
||||
token = token,
|
||||
)
|
||||
|
||||
print(tokenizer)
|
||||
print(tokenizer.chat_template)
|
||||
|
||||
model, tokenizer = patch_tokenizer(model, tokenizer)
|
||||
print(tokenizer)
|
||||
print(tokenizer.chat_template)
|
||||
model = model_patcher.post_patch(model)
|
||||
|
||||
# Patch up QKV / O and MLP
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ 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
|
||||
|
||||
import os
|
||||
|
||||
# https://github.com/huggingface/transformers/pull/26037 allows 4 bit loading!
|
||||
major, minor = transformers_version.split(".")[:2]
|
||||
|
|
@ -79,14 +79,12 @@ class FastLanguageModel(FastLlamaModel):
|
|||
):
|
||||
old_model_name = model_name
|
||||
model_name = _get_model_name(model_name, load_in_4bit)
|
||||
print(model_name)
|
||||
|
||||
# First check if it's a normal model via AutoConfig
|
||||
is_peft = False
|
||||
try:
|
||||
model_config = AutoConfig.from_pretrained(model_name, token = token)
|
||||
is_peft = False
|
||||
print(model_config)
|
||||
except:
|
||||
try:
|
||||
# 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_config = AutoConfig.from_pretrained(model_name, token = token)
|
||||
is_peft = True
|
||||
print(model_config)
|
||||
pass
|
||||
|
||||
model_type = model_config.model_type
|
||||
|
|
@ -121,7 +118,16 @@ class FastLanguageModel(FastLlamaModel):
|
|||
)
|
||||
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_name = model_name,
|
||||
max_seq_length = max_seq_length,
|
||||
|
|
@ -132,6 +138,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
rope_scaling = rope_scaling,
|
||||
fix_tokenizer = fix_tokenizer,
|
||||
model_patcher = dispatch_model,
|
||||
tokenizer_name = tokenizer_name,
|
||||
*args, **kwargs,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -294,6 +294,7 @@ class FastMistralModel(FastLlamaModel):
|
|||
rope_scaling = None, # Mistral does not support RoPE scaling
|
||||
fix_tokenizer = True,
|
||||
model_patcher = None,
|
||||
tokenizer_name = None,
|
||||
**kwargs,
|
||||
):
|
||||
if model_patcher is None: model_patcher = FastMistralModel
|
||||
|
|
@ -354,8 +355,11 @@ class FastMistralModel(FastLlamaModel):
|
|||
# rope_scaling = rope_scaling,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Counteract saved tokenizers
|
||||
tokenizer_name = model_name if tokenizer_name is None else tokenizer_name
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_name,
|
||||
tokenizer_name,
|
||||
model_max_length = max_position_embeddings,
|
||||
padding_side = "right",
|
||||
token = token,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue