diff --git a/unsloth/chat_templates.py b/unsloth/chat_templates.py index b019111057..7c3c74f473 100644 --- a/unsloth/chat_templates.py +++ b/unsloth/chat_templates.py @@ -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 <> system messages. +llama_template = \ + "{% if messages[0]['role'] == 'system' %}"\ + "{% if messages[1]['role'] == 'user' %}"\ + "{{ bos_token + '[INST] <>\n' + messages[0]['content'] + '\n<>\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