patch tokenizer

This commit is contained in:
Daniel Han-Chen 2024-02-14 17:56:14 +11:00
commit efbb1e6049
2 changed files with 8 additions and 3 deletions

View file

@ -20,6 +20,7 @@ __all__ = [
from transformers import StoppingCriteria, StoppingCriteriaList
from torch import LongTensor, FloatTensor
from transformers.models.llama.modeling_llama import logger
from .models._utils import patch_tokenizer
CHAT_TEMPLATES = {}
@ -263,6 +264,7 @@ def get_chat_template(
.replace("'user'", "'" + mapping["user"] + "'")\
.replace("'assistant'", "'" + mapping["assistant"] + "'")
_, tokenizer = patch_tokenizer(model = None, tokenizer = tokenizer)
tokenizer.chat_template = chat_template
#stopping_criteria = create_stopping_criteria(tokenizer, stop_word)

View file

@ -117,21 +117,24 @@ pass
def patch_tokenizer(model, tokenizer):
model.config.update({"unsloth_version" : __version__})
if model is not None:
model.config.update({"unsloth_version" : __version__})
if not hasattr(tokenizer, "pad_token") or tokenizer.pad_token is None:
# Fixes https://github.com/unslothai/unsloth/issues/5
if hasattr(tokenizer, "unk_token"):
tokenizer.add_special_tokens({"pad_token" : tokenizer.unk_token})
tokenizer.pad_token = tokenizer.unk_token
else:
name = model.config._name_or_path if model is not None else "Model"
logger.warning_one(
f"{model.config._name_or_path} does not have a padding or unknown token!\n"\
f"{name} does not have a padding or unknown token!\n"\
f"Will use the EOS token of id {tokenizer.eos_token_id} as padding."
)
assert(hasattr(tokenizer, "eos_token"))
tokenizer.add_special_tokens({"pad_token" : tokenizer.eos_token})
tokenizer.pad_token = tokenizer.eos_token
config = model.config.update({"pad_token_id" : tokenizer.eos_token_id})
if model is not None:
config = model.config.update({"pad_token_id" : tokenizer.eos_token_id})
pass
return model, tokenizer
pass