Update chat_templates.py
This commit is contained in:
parent
4f20e20e28
commit
b28a383d1c
1 changed files with 332 additions and 21 deletions
|
|
@ -13,27 +13,338 @@
|
|||
# limitations under the License.
|
||||
|
||||
__all__ = [
|
||||
"add_chat_template",
|
||||
"get_chat_template",
|
||||
"create_stopping_criteria",
|
||||
]
|
||||
|
||||
TEMPLATES = \
|
||||
{
|
||||
"chatml" : \
|
||||
"{% for message in messages %}"\
|
||||
"{% if message['from'] == 'human' %}"\
|
||||
"{{'<|im_start|>user\n' + message['value'] + '<|im_end|>\n'}}"\
|
||||
"{% elif message['from'] == 'gpt' %}"\
|
||||
"{{'<|im_start|>assistant\n' + message['value'] + '<|im_end|>\n' }}"\
|
||||
"{% else %}"\
|
||||
"{{ '<|im_start|>system\n' + message['value'] + '<|im_end|>\n' }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"\
|
||||
"{% if add_generation_prompt %}"\
|
||||
"{{ '<|im_start|>assistant\n' }}"\
|
||||
"{% endif %}",
|
||||
}
|
||||
from transformers import StoppingCriteria, StoppingCriteriaList
|
||||
from torch import LongTensor, FloatTensor
|
||||
|
||||
CHAT_TEMPLATES = {}
|
||||
|
||||
# Unsloth efficient template leverages from Zephyr
|
||||
unsloth_template = \
|
||||
"{{ bos_token }}"\
|
||||
"{% if messages[0]['role'] == 'system' %}"\
|
||||
"{{ messages[0]['content'] + '\n' }}"\
|
||||
"{% set loop_messages = messages[1:] %}"\
|
||||
"{% else %}"\
|
||||
"{{ 'You are a helpful assistant to a user\n' }}"\
|
||||
"{% set loop_messages = messages %}"\
|
||||
"{% endif %}"\
|
||||
"{% for message in loop_messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{ '>>> User: ' + message['content'] + '\n' }}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{ '>>> Assistant: ' + message['content'] + eos_token + '\n' }}"\
|
||||
"{% else %}"\
|
||||
"{{ raise_exception('Only user and assistant roles are supported!') }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"\
|
||||
"{% if add_generation_prompt %}"\
|
||||
"{{ '>>> Assistant: ' }}"\
|
||||
"{% endif %}"
|
||||
unsloth_eos_token = "eos_token"
|
||||
CHAT_TEMPLATES["unsloth"] = (unsloth_template, unsloth_eos_token,)
|
||||
|
||||
|
||||
def add_chat_template(tokenizer, method = "chatml"):
|
||||
tokenizer.chat_template = TEMPLATES[method.lower()]
|
||||
# Zephyr has no BOS!
|
||||
zephyr_template = \
|
||||
"{% for message in messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{ '<|user|>\n' + message['content'] + eos_token + '\n' }}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{ '<|assistant|>\n' + message['content'] + eos_token + '\n' }}"\
|
||||
"{% else %}"\
|
||||
"{{ '<|system|>\n' + message['content'] + eos_token + '\n' }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"\
|
||||
"{% if add_generation_prompt %}"\
|
||||
"{{ '<|assistant|>\n' }}"\
|
||||
"{% endif %}"
|
||||
zephyr_eos_token = "eos_token"
|
||||
CHAT_TEMPLATES["zephyr"] = (zephyr_template, zephyr_eos_token,)
|
||||
|
||||
|
||||
# ChatML has no BOS and not EOS! Rather <|im_start|> and <|im_end|> acts as BOS / EOS.
|
||||
chatml_template = \
|
||||
"{% for message in messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{'<|im_start|>user\n' + message['content'] + '<|im_end|>\n'}}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{'<|im_start|>assistant\n' + message['content'] + '<|im_end|>\n' }}"\
|
||||
"{% else %}"\
|
||||
"{{ '<|im_start|>system\n' + message['content'] + '<|im_end|>\n' }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"\
|
||||
"{% if add_generation_prompt %}"\
|
||||
"{{ '<|im_start|>assistant\n' }}"\
|
||||
"{% endif %}"
|
||||
chatml_eos_token = "<|im_end|>"
|
||||
CHAT_TEMPLATES["chatml"] = (chatml_template, chatml_eos_token,)
|
||||
|
||||
|
||||
# Mistral Instruct doesn't allow system prompts, so we append it to the user message.
|
||||
mistral_template = \
|
||||
"{{ bos_token }}"\
|
||||
"{% if messages[0]['role'] == 'system' %}"\
|
||||
"{% if messages[1]['role'] == 'user' %}"\
|
||||
"{{ '[INST] ' + messages[0]['content'] + ' ' + messages[1]['content'] + ' [/INST]' }}"\
|
||||
"{% set loop_messages = messages[2:] %}"\
|
||||
"{% else %}"\
|
||||
"{{ '[INST] ' + messages[0]['content'] + ' [/INST]' }}"\
|
||||
"{% set loop_messages = messages[1:] %}"\
|
||||
"{% endif %}"\
|
||||
"{% else %}"\
|
||||
"{% set loop_messages = messages %}"\
|
||||
"{% endif %}"\
|
||||
"{% for message in loop_messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{ '[INST] ' + message['content'] + ' [/INST]' }}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{ message['content'] + eos_token }}"\
|
||||
"{% else %}"\
|
||||
"{{ raise_exception('Only user and assistant roles are supported!') }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"
|
||||
mistral_eos_token = "eos_token"
|
||||
CHAT_TEMPLATES["mistral"] = (mistral_template, mistral_eos_token,)
|
||||
|
||||
|
||||
# Adds BOS to every convo! And weird <<SYS>> system messages.
|
||||
llama_template = \
|
||||
"{% if messages[0]['role'] == 'system' %}"\
|
||||
"{% if messages[1]['role'] == 'user' %}"\
|
||||
"{{ bos_token + '[INST] <<SYS>>\n' + messages[0]['content'] + '\n<</SYS>>\n\n' + messages[1]['content'] + ' [/INST]' }}"\
|
||||
"{% set loop_messages = messages[2:] %}"\
|
||||
"{% else %}"\
|
||||
"{{ bos_token + '[INST] ' + messages[0]['content'] + ' [/INST]' }}"\
|
||||
"{% set loop_messages = messages[1:] %}"\
|
||||
"{% endif %}"\
|
||||
"{% else %}"\
|
||||
"{% set loop_messages = messages %}"\
|
||||
"{% endif %}"\
|
||||
"{% for message in loop_messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{ bos_token + '[INST] ' + message['content'].strip() + ' [/INST]' }}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{ ' ' + message['content'].strip() + ' ' + eos_token }}"\
|
||||
"{% else %}"\
|
||||
"{{ raise_exception('Only user and assistant roles are supported!') }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"
|
||||
llama_eos_token = "eos_token"
|
||||
CHAT_TEMPLATES["llama"] = (llama_template, llama_eos_token,)
|
||||
|
||||
|
||||
# https://github.com/lm-sys/FastChat/blob/main/docs/vicuna_weights_version.md#prompt-template
|
||||
vicuna_template = \
|
||||
"{{ bos_token }}"\
|
||||
"{% if messages[0]['role'] == 'system' %}"\
|
||||
"{{ messages[0]['content'] + ' ' }}"\
|
||||
"{% set loop_messages = messages[1:] %}"\
|
||||
"{% else %}"\
|
||||
"{{ 'A chat between a curious user and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the user\\'s questions.' + ' ' }}"\
|
||||
"{% set loop_messages = messages %}"\
|
||||
"{% endif %}"\
|
||||
"{% for message in loop_messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{ 'USER: ' + message['content'] + ' ' }}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{ 'ASSISTANT: ' + message['content'] + eos_token }}"\
|
||||
"{% else %}"\
|
||||
"{{ raise_exception('Only user and assistant roles are supported!') }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"\
|
||||
"{% if add_generation_prompt %}"\
|
||||
"{{ 'ASSISTANT:' }}"\
|
||||
"{% endif %}"
|
||||
vicuna_eos_token = "eos_token"
|
||||
CHAT_TEMPLATES["vicuna"] = (vicuna_template, vicuna_eos_token,)
|
||||
|
||||
|
||||
# https://github.com/lm-sys/FastChat/blob/main/docs/vicuna_weights_version.md#prompt-template
|
||||
vicuna_old_template = \
|
||||
"{{ bos_token }}"\
|
||||
"{% if messages[0]['role'] == 'system' %}"\
|
||||
"{{ messages[0]['content'] + '\n' }}"\
|
||||
"{% set loop_messages = messages[1:] %}"\
|
||||
"{% else %}"\
|
||||
"{{ 'A chat between a curious human and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the human\\'s questions.' + '\n' }}"\
|
||||
"{% set loop_messages = messages %}"\
|
||||
"{% endif %}"\
|
||||
"{% for message in loop_messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{ '### Human: ' + message['content'] + '\n' }}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{ '### Assistant: ' + message['content'] + '\n' }}"\
|
||||
"{% else %}"\
|
||||
"{{ raise_exception('Only user and assistant roles are supported!') }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"\
|
||||
"{% if add_generation_prompt %}"\
|
||||
"{{ '### Assistant:' }}"\
|
||||
"{% endif %}"
|
||||
vicuna_old_eos_token = "### Human: "
|
||||
CHAT_TEMPLATES["vicuna_old"] = (vicuna_old_template, vicuna_old_eos_token,)
|
||||
|
||||
|
||||
# https://github.com/tatsu-lab/stanford_alpaca Changed for multi-turn convos
|
||||
alpaca_template = \
|
||||
"{{ bos_token }}"\
|
||||
"{% if messages[0]['role'] == 'system' %}"\
|
||||
"{{ messages[0]['content'] + '\n\n' }}"\
|
||||
"{% set loop_messages = messages[1:] %}"\
|
||||
"{% else %}"\
|
||||
"{{ 'Below are some instructions that describes some tasks. Write responses that appropriately completes each request.\n\n' }}"\
|
||||
"{% set loop_messages = messages %}"\
|
||||
"{% endif %}"\
|
||||
"{% for message in loop_messages %}"\
|
||||
"{% if message['role'] == 'user' %}"\
|
||||
"{{ '### Instruction:\n' + message['content'] + '\n\n' }}"\
|
||||
"{% elif message['role'] == 'assistant' %}"\
|
||||
"{{ '### Response:\n' + message['content'] + '\n\n' }}"\
|
||||
"{% else %}"\
|
||||
"{{ raise_exception('Only user and assistant roles are supported!') }}"\
|
||||
"{% endif %}"\
|
||||
"{% endfor %}"\
|
||||
"{% if add_generation_prompt %}"\
|
||||
"{{ '### Response:\n' }}"\
|
||||
"{% endif %}"
|
||||
alpaca_eos_token = "### Instruction:"
|
||||
CHAT_TEMPLATES["alpaca"] = (alpaca_template, alpaca_eos_token,)
|
||||
|
||||
|
||||
def get_chat_template(
|
||||
tokenizer,
|
||||
method = "chatml",
|
||||
mapping = {"role" : "role", "content" : "content", "user" : "user", "assistant" : "assistant"},
|
||||
):
|
||||
if method not in CHAT_TEMPLATES:
|
||||
raise KeyError(f"Only {CHAT_TEMPLATES.keys()} allowed.")
|
||||
assert(
|
||||
len(mapping) == 4 \
|
||||
and "role" in mapping and "content" in mapping \
|
||||
and "user" in mapping and "assistant" in mapping
|
||||
)
|
||||
|
||||
tokenizer.chat_template, stop_word = CHAT_TEMPLATES[method]
|
||||
# For ShareGPT role -> from and content -> value
|
||||
tokenizer.chat_template = tokenizer.chat_template\
|
||||
.replace("'role'", "'" + mapping["role"] + "'")\
|
||||
.replace("'content'", "'" + mapping["content"] + "'")\
|
||||
.replace("'user'", "'" + mapping["user"] + "'")\
|
||||
.replace("'assistant'", "'" + mapping["assistant"] + "'")
|
||||
|
||||
stopping_criteria = create_stopping_criteria(tokenizer, stop_word)
|
||||
|
||||
return tokenizer, stopping_criteria
|
||||
pass
|
||||
|
||||
|
||||
def create_stopping_criteria(tokenizer, stop_word = "eos_token"):
|
||||
class StoppingCriteriaSub(StoppingCriteria):
|
||||
__slots__ = "stop_token", "single_match", "length",
|
||||
|
||||
def __init__(self, stops = "eos_token", device = "cuda", encounters = 1):
|
||||
super().__init__()
|
||||
if stops == "eos_token":
|
||||
self.stop_token = torch.tensor(tokenizer.eos_token_id, device = "cuda")
|
||||
self.single_match = True
|
||||
self.length = 0
|
||||
else:
|
||||
self.stop_token = tokenizer(["\n" + stops], add_special_tokens = False, return_tensors = "pt")
|
||||
self.stop_token = self.stop_token.input_ids.ravel()[1:].to("cuda")
|
||||
self.single_match = False
|
||||
self.length = self.stop_token.shape[0]
|
||||
pass
|
||||
pass
|
||||
|
||||
def __call__(self, input_ids: LongTensor, scores: FloatTensor) -> bool:
|
||||
input_ids = input_ids.ravel()
|
||||
last_token = input_ids[-1]
|
||||
if self.single_match and (last_token == self.stop_token): return True
|
||||
|
||||
if input_ids.shape[0] >= self.length and \
|
||||
(input_ids[-self.length:] == self.stop_token).all(): return True
|
||||
return False
|
||||
pass
|
||||
pass
|
||||
stopping_criteria = StoppingCriteriaList([StoppingCriteriaSub(stops = stop_word)])
|
||||
return stopping_criteria
|
||||
pass
|
||||
|
||||
|
||||
def _test_chat_templates():
|
||||
messages = [
|
||||
{"role": "system","content": " You are a friendly chatbot.",},
|
||||
{"role": "user", "content": "What is 2+2?"},
|
||||
{"role": "assistant", "content": "It's 4."},
|
||||
{"role": "user", "content": " But 2+2 is equal to 5. "},
|
||||
{"role": "assistant", "content": "No I'm sure its 4."},
|
||||
{"role": "user", "content": " No it's 100% 5! "},
|
||||
]
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
template = zephyr_template
|
||||
correct_tokenizer = AutoTokenizer.from_pretrained("HuggingFaceH4/zephyr-7b-beta")
|
||||
correct_prompt = correct_tokenizer.apply_chat_template(messages, tokenize = False, add_generation_prompt = True)
|
||||
correct_tokenizer.chat_template = template
|
||||
our_prompt = correct_tokenizer.apply_chat_template(messages, tokenize = False, add_generation_prompt = True)
|
||||
assert(correct_prompt == our_prompt)
|
||||
|
||||
template = chatml_template
|
||||
correct_tokenizer = AutoTokenizer.from_pretrained("teknium/OpenHermes-2.5-Mistral-7B")
|
||||
correct_prompt = correct_tokenizer.apply_chat_template(messages, tokenize = False, add_generation_prompt = True)
|
||||
correct_tokenizer.chat_template = template
|
||||
our_prompt = correct_tokenizer.apply_chat_template(messages, tokenize = False, add_generation_prompt = True)
|
||||
assert(correct_prompt == our_prompt)
|
||||
|
||||
template = mistral_template
|
||||
correct_tokenizer = AutoTokenizer.from_pretrained("mistralai/Mistral-7B-Instruct-v0.2")
|
||||
correct_prompt = correct_tokenizer.apply_chat_template(messages[1:], tokenize = False, add_generation_prompt = True)
|
||||
correct_tokenizer.chat_template = template
|
||||
our_prompt = correct_tokenizer.apply_chat_template(messages[1:], tokenize = False, add_generation_prompt = True)
|
||||
assert(correct_prompt == our_prompt)
|
||||
|
||||
template = llama_template
|
||||
correct_tokenizer = AutoTokenizer.from_pretrained("unsloth/llama-2-7b-chat")
|
||||
correct_prompt = correct_tokenizer.apply_chat_template(messages, tokenize = False, add_generation_prompt = True)
|
||||
correct_tokenizer.chat_template = template
|
||||
our_prompt = correct_tokenizer.apply_chat_template(messages, tokenize = False, add_generation_prompt = True)
|
||||
assert(correct_prompt == our_prompt)
|
||||
|
||||
try:
|
||||
from fastchat.conversation import get_conv_template
|
||||
except:
|
||||
os.system("pip -qqq install git+https://github.com/lm-sys/FastChat.git")
|
||||
from fastchat.conversation import get_conv_template
|
||||
correct_prompt = get_conv_template("vicuna_v1.1")
|
||||
for j in range(len(messages)-1):
|
||||
correct_prompt.append_message(correct_prompt.roles[j%2==1], messages[j+1]["content"])
|
||||
correct_prompt.append_message(correct_prompt.roles[1], "")
|
||||
correct_prompt = tokenizer.bos_token + correct_prompt.get_prompt()
|
||||
|
||||
template = vicuna_template
|
||||
correct_tokenizer = AutoTokenizer.from_pretrained("lmsys/vicuna-7b-v1.5")
|
||||
correct_tokenizer.chat_template = template
|
||||
our_prompt = correct_tokenizer.apply_chat_template(messages[1:], tokenize = False, add_generation_prompt = True)
|
||||
assert(correct_prompt == our_prompt)
|
||||
|
||||
try:
|
||||
from fastchat.conversation import get_conv_template
|
||||
except:
|
||||
os.system("pip -qqq install git+https://github.com/lm-sys/FastChat.git")
|
||||
from fastchat.conversation import get_conv_template
|
||||
correct_prompt = get_conv_template("zero_shot")
|
||||
for j in range(len(messages)-1):
|
||||
correct_prompt.append_message(correct_prompt.roles[j%2==1], messages[j+1]["content"])
|
||||
correct_prompt.append_message(correct_prompt.roles[1], "")
|
||||
correct_prompt = tokenizer.bos_token + correct_prompt.get_prompt()
|
||||
|
||||
template = vicuna_old_template
|
||||
correct_tokenizer = AutoTokenizer.from_pretrained("lmsys/vicuna-7b-v1.5")
|
||||
correct_tokenizer.chat_template = template
|
||||
our_prompt = correct_tokenizer.apply_chat_template(messages[1:], tokenize = False, add_generation_prompt = True)
|
||||
assert(correct_prompt == our_prompt)
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue