hymenjj/llama-cpp-python-prebuilt
0
1from __future__ import annotations2 3import os4import sys5import json6import ctypes7import dataclasses8import random9import string10 11from datetime import datetime12from contextlib import ExitStack13from typing import (14 Any,15 Dict,16 Iterator,17 List,18 Literal,19 Optional,20 Tuple,21 Union,22 Protocol,23 cast,24)25 26import jinja227from jinja2.sandbox import ImmutableSandboxedEnvironment28 29import numpy as np30import numpy.typing as npt31 32import llama_cpp.llama_cpp as llama_cpp33import llama_cpp.llama as llama34import llama_cpp.llama_types as llama_types35import llama_cpp.llama_grammar as llama_grammar36 37from ._logger import logger38from ._utils import suppress_stdout_stderr, Singleton39 40### Common Chat Templates and Special Tokens ###41 42# Source: https://huggingface.co/teknium/OpenHermes-2.5-Mistral-7B/blob/main/tokenizer_config.json43CHATML_CHAT_TEMPLATE = "{% for message in messages %}{{'<|im_start|>' + message['role'] + '\n' + message['content'] + '<|im_end|>' + '\n'}}{% endfor %}{% if add_generation_prompt %}{{ '<|im_start|>assistant\n' }}{% endif %}"44CHATML_BOS_TOKEN = "<s>"45CHATML_EOS_TOKEN = "<|im_end|>"46 47# Source: https://huggingface.co/mistralai/Mistral-7B-Instruct-v0.1/blob/main/tokenizer_config.json48MISTRAL_INSTRUCT_CHAT_TEMPLATE = "{{ bos_token }}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% 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 %}"49MISTRAL_INSTRUCT_BOS_TOKEN = "<s>"50MISTRAL_INSTRUCT_EOS_TOKEN = "</s>"51 52# Source: https://huggingface.co/mistralai/Mixtral-8x7B-Instruct-v0.1/blob/main/tokenizer_config.json53MIXTRAL_INSTRUCT_CHAT_TEMPLATE = "{{ bos_token }}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% 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 %}"54 55# Source: https://huggingface.co/meta-llama/Meta-Llama-3-8B-Instruct/blob/main/tokenizer_config.json56LLAMA3_INSTRUCT_CHAT_TEMPLATE = "{% set loop_messages = messages %}{% for message in loop_messages %}{% set content = '<|start_header_id|>' + message['role'] + '<|end_header_id|>\n\n'+ message['content'] | trim + '<|eot_id|>' %}{% if loop.index0 == 0 %}{% set content = bos_token + content %}{% endif %}{{ content }}{% endfor %}{% if add_generation_prompt %}{{ '<|start_header_id|>assistant<|end_header_id|>\n\n' }}{% endif %}"57 58### Chat Completion Handler ###59 60 61class LlamaChatCompletionHandler(Protocol):62 """Base Protocol for a llama chat completion handler.63 64 Very generic protocol that can be used to implement any chat format.65 The only hard requirement is that it must return a ChatCompletion when66 stream=False and an iterator of ChatCompletionChunks when stream=True."""67 68 def __call__(69 self,70 *,71 # llama.cpp instance72 llama: llama.Llama,73 # openai api parameters74 messages: List[llama_types.ChatCompletionRequestMessage],75 functions: Optional[List[llama_types.ChatCompletionFunction]] = None,76 function_call: Optional[llama_types.ChatCompletionRequestFunctionCall] = None,77 tools: Optional[List[llama_types.ChatCompletionTool]] = None,78 tool_choice: Optional[llama_types.ChatCompletionToolChoiceOption] = None,79 temperature: float = 0.2,80 top_p: float = 0.95,81 top_k: int = 40,82 stream: bool = False,83 stop: Optional[Union[str, List[str]]] = [],84 seed: Optional[int] = None,85 response_format: Optional[86 llama_types.ChatCompletionRequestResponseFormat87 ] = None,88 max_tokens: Optional[int] = None,89 presence_penalty: float = 0.0,90 frequency_penalty: float = 0.0,91 repeat_penalty: float = 1.1,92 model: Optional[str] = None,93 logit_bias: Optional[Dict[str, float]] = None,94 # llama.cpp parameters95 min_p: float = 0.05,96 typical_p: float = 1.0,97 tfs_z: float = 1.0,98 mirostat_mode: int = 0,99 mirostat_tau: float = 5.0,100 mirostat_eta: float = 0.1,101 logits_processor: Optional[llama.LogitsProcessorList] = None,102 grammar: Optional[llama.LlamaGrammar] = None,103 logprobs: Optional[bool] = None,104 top_logprobs: Optional[int] = None,105 **kwargs, # type: ignore106 ) -> Union[107 llama_types.CreateChatCompletionResponse,108 Iterator[llama_types.CreateChatCompletionStreamResponse],109 ]: ...110 111 112class LlamaChatCompletionHandlerNotFoundException(Exception):113 pass114 115 116class LlamaChatCompletionHandlerRegistry(Singleton):117 _chat_handlers: Dict[str, LlamaChatCompletionHandler] = {}118 119 def register_chat_completion_handler(120 self,121 name: str,122 chat_handler: LlamaChatCompletionHandler,123 overwrite: bool = False,124 ):125 if not overwrite and name in self._chat_handlers:126 raise ValueError(127 f"Formatter with name '{name}' is already registered. Use `overwrite=True` to overwrite it."128 )129 self._chat_handlers[name] = chat_handler130 131 def unregister_chat_handler(self, name: str):132 if name in self._chat_handlers:133 del self._chat_handlers[name]134 else:135 raise ValueError(f"No formatter registered under the name '{name}'.")136 137 def get_chat_completion_handler_by_name(138 self, name: str139 ) -> LlamaChatCompletionHandler:140 try:141 chat_handler = self._chat_handlers[name]142 return chat_handler143 except KeyError:144 raise LlamaChatCompletionHandlerNotFoundException(145 f"Invalid chat handler: {name} (valid formats: {list(self._chat_handlers.keys())})"146 )147 148 149def get_chat_completion_handler(name: str) -> LlamaChatCompletionHandler:150 return LlamaChatCompletionHandlerRegistry().get_chat_completion_handler_by_name(151 name152 )153 154 155def register_chat_completion_handler(name: str):156 def decorator(f: LlamaChatCompletionHandler):157 LlamaChatCompletionHandlerRegistry().register_chat_completion_handler(name, f)158 return f159 160 return decorator161 162 163### Chat Formatter ###164 165 166@dataclasses.dataclass167class ChatFormatterResponse:168 """Dataclass that stores completion parameters for a given chat format and169 create_chat_completion request.170 171 prompt contains the formatted prompt generated from the chat format and messages.172 stop contains the stop token or list of stop tokens to use for the chat format."""173 174 prompt: str175 stop: Optional[Union[str, List[str]]] = None176 stopping_criteria: Optional[llama.StoppingCriteriaList] = None177 added_special: bool = False178 179 180class ChatFormatter(Protocol):181 """Base Protocol for a chat formatter. A chat formatter is a function that182 takes a list of messages and returns a chat format response which can be used183 to generate a completion. The response can also include a stop token or list184 of stop tokens to use for the completion."""185 186 def __call__(187 self,188 *,189 messages: List[llama_types.ChatCompletionRequestMessage],190 **kwargs: Any,191 ) -> ChatFormatterResponse: ...192 193 194class Jinja2ChatFormatter(ChatFormatter):195 def __init__(196 self,197 template: str,198 eos_token: str,199 bos_token: str,200 add_generation_prompt: bool = True,201 stop_token_ids: Optional[List[int]] = None,202 ):203 """A chat formatter that uses jinja2 templates to format the prompt."""204 self.template = template205 self.eos_token = eos_token206 self.bos_token = bos_token207 self.add_generation_prompt = add_generation_prompt208 self.stop_token_ids = (209 set(stop_token_ids) if stop_token_ids is not None else None210 )211 212 self._environment = ImmutableSandboxedEnvironment(213 loader=jinja2.BaseLoader(),214 trim_blocks=True,215 lstrip_blocks=True,216 ).from_string(self.template)217 218 @staticmethod219 def strftime_now(f: str) -> str:220 return datetime.now().strftime(f)221 222 def __call__(223 self,224 *,225 messages: List[llama_types.ChatCompletionRequestMessage],226 functions: Optional[List[llama_types.ChatCompletionFunction]] = None,227 function_call: Optional[llama_types.ChatCompletionRequestFunctionCall] = None,228 tools: Optional[List[llama_types.ChatCompletionTool]] = None,229 tool_choice: Optional[llama_types.ChatCompletionToolChoiceOption] = None,230 **kwargs: Any,231 ) -> ChatFormatterResponse:232 def raise_exception(message: str):233 raise ValueError(message)234 235 prompt = self._environment.render(236 messages=messages,237 eos_token=self.eos_token,238 bos_token=self.bos_token,239 raise_exception=raise_exception,240 add_generation_prompt=self.add_generation_prompt,241 functions=functions,242 function_call=function_call,243 tools=tools,244 tool_choice=tool_choice,245 strftime_now=self.strftime_now,246 )247 248 stopping_criteria = None249 if self.stop_token_ids is not None:250 251 def stop_on_last_token(252 tokens: npt.NDArray[np.intc], logits: npt.NDArray[np.single]253 ) -> bool:254 return tokens[-1] in self.stop_token_ids255 256 stopping_criteria = llama.StoppingCriteriaList([stop_on_last_token])257 258 return ChatFormatterResponse(259 prompt=prompt,260 stop=[self.eos_token],261 stopping_criteria=stopping_criteria,262 added_special=True,263 )264 265 def to_chat_handler(self) -> LlamaChatCompletionHandler:266 return chat_formatter_to_chat_completion_handler(self)267 268 269def _convert_text_completion_logprobs_to_chat(270 logprobs: Optional[llama_types.CompletionLogprobs],271) -> llama_types.ChatCompletionLogprobs:272 if logprobs is None:273 return None274 275 return {276 "content": [277 {278 "token": token,279 "bytes": None,280 "logprob": logprob,281 "top_logprobs": [282 {283 "token": top_token,284 "logprob": top_logprob,285 "bytes": None,286 }287 for top_token, top_logprob in top_logprobs.items()288 ],289 } for (token, logprob, top_logprobs) in zip(logprobs["tokens"], logprobs["token_logprobs"], logprobs["top_logprobs"])290 ],291 "refusal": None,292 }293 294def _convert_text_completion_to_chat(295 completion: llama_types.Completion,296) -> llama_types.ChatCompletion:297 assert "usage" in completion298 return {299 "id": "chat" + completion["id"],300 "object": "chat.completion",301 "created": completion["created"],302 "model": completion["model"],303 "choices": [304 {305 "index": 0,306 "message": {307 "role": "assistant",308 "content": completion["choices"][0]["text"],309 },310 "logprobs": _convert_text_completion_logprobs_to_chat(completion["choices"][0]["logprobs"]),311 "finish_reason": completion["choices"][0]["finish_reason"],312 }313 ],314 "usage": completion["usage"],315 }316 317 318def _convert_text_completion_chunks_to_chat(319 chunks: Iterator[llama_types.CreateCompletionStreamResponse],320) -> Iterator[llama_types.ChatCompletionChunk]:321 for i, chunk in enumerate(chunks):322 if i == 0:323 yield {324 "id": "chat" + chunk["id"],325 "model": chunk["model"],326 "created": chunk["created"],327 "object": "chat.completion.chunk",328 "choices": [329 {330 "index": 0,331 "delta": {332 "role": "assistant",333 },334 "logprobs": None,335 "finish_reason": None,336 }337 ],338 }339 yield {340 "id": "chat" + chunk["id"],341 "model": chunk["model"],342 "created": chunk["created"],343 "object": "chat.completion.chunk",344 "choices": [345 {346 "index": 0,347 "delta": (348 {349 "content": chunk["choices"][0]["text"],350 }351 if chunk["choices"][0]["finish_reason"] is None352 else {}353 ),354 "logprobs": _convert_text_completion_logprobs_to_chat(chunk["choices"][0]["logprobs"]),355 "finish_reason": chunk["choices"][0]["finish_reason"],356 }357 ],358 }359 360 361def _convert_completion_to_chat(362 completion_or_chunks: Union[363 llama_types.CreateCompletionResponse,364 Iterator[llama_types.CreateCompletionStreamResponse],365 ],366 stream: bool = False,367) -> Union[368 llama_types.CreateChatCompletionResponse, Iterator[llama_types.ChatCompletionChunk]369]:370 if stream:371 chunks: Iterator[llama_types.CreateCompletionStreamResponse] = completion_or_chunks # type: ignore372 return _convert_text_completion_chunks_to_chat(chunks)373 else:374 completion: llama_types.Completion = completion_or_chunks # type: ignore375 return _convert_text_completion_to_chat(completion)376 377 378def _convert_completion_to_chat_function(379 tool_name: str,380 completion_or_chunks: Union[381 llama_types.CreateCompletionResponse,382 Iterator[llama_types.CreateCompletionStreamResponse],383 ],384 stream: bool,385):386 if not stream:387 completion: llama_types.CreateCompletionResponse = completion_or_chunks # type: ignore388 assert "usage" in completion389 tool_id = "call_" + "_0_" + tool_name + "_" + completion["id"]390 # TODO: Fix for legacy function calls391 chat_completion: llama_types.CreateChatCompletionResponse = {392 "id": "chat" + completion["id"],393 "object": "chat.completion",394 "created": completion["created"],395 "model": completion["model"],396 "choices": [397 {398 "index": 0,399 "message": {400 "role": "assistant",401 "content": None,402 "function_call": {403 "name": tool_name,404 "arguments": completion["choices"][0]["text"],405 },406 "tool_calls": [407 {408 "id": tool_id,409 "type": "function",410 "function": {411 "name": tool_name,412 "arguments": completion["choices"][0]["text"],413 },414 }415 ],416 },417 "logprobs": _convert_text_completion_logprobs_to_chat(completion["choices"][0]["logprobs"]),418 "finish_reason": "tool_calls",419 }420 ],421 "usage": completion["usage"],422 }423 return chat_completion424 else:425 chunks: Iterator[llama_types.CreateCompletionStreamResponse] = completion_or_chunks # type: ignore426 427 def _stream_response_to_function_stream(428 chunks: Iterator[llama_types.CreateCompletionStreamResponse],429 ) -> Iterator[llama_types.CreateChatCompletionStreamResponse]:430 # blank first message431 first = True432 id_ = None433 created = None434 model = None435 tool_id = None436 for chunk in chunks:437 if first:438 id_ = "chat" + chunk["id"]439 created = chunk["created"]440 model = chunk["model"]441 tool_id = "call_" + "_0_" + tool_name + "_" + chunk["id"]442 yield {443 "id": id_,444 "object": "chat.completion.chunk",445 "created": created,446 "model": model,447 "choices": [448 {449 "index": 0,450 "finish_reason": None,451 "logprobs": None,452 "delta": {453 "role": "assistant",454 "content": None,455 "function_call": None,456 "tool_calls": None,457 },458 }459 ],460 }461 yield {462 "id": "chat" + chunk["id"],463 "object": "chat.completion.chunk",464 "created": chunk["created"],465 "model": chunk["model"],466 "choices": [467 {468 "index": 0,469 "finish_reason": None,470 "logprobs": _convert_text_completion_logprobs_to_chat(chunk["choices"][0]["logprobs"]),471 "delta": {472 "role": None,473 "content": None,474 "function_call": {475 "name": tool_name,476 "arguments": chunk["choices"][0]["text"],477 },478 "tool_calls": [479 {480 "index": 0,481 "id": tool_id,482 "type": "function",483 "function": {484 "name": tool_name,485 "arguments": chunk["choices"][0][486 "text"487 ],488 },489 }490 ],491 },492 }493 ],494 }495 first = False496 continue497 assert tool_id is not None498 yield {499 "id": "chat" + chunk["id"],500 "object": "chat.completion.chunk",501 "created": chunk["created"],502 "model": chunk["model"],503 "choices": [504 {505 "index": 0,506 "finish_reason": None,507 "logprobs": _convert_text_completion_logprobs_to_chat(chunk["choices"][0]["logprobs"]),508 "delta": {509 "role": None,510 "content": None,511 "function_call": {512 "name": tool_name,513 "arguments": chunk["choices"][0]["text"],514 },515 "tool_calls": [516 {517 "index": 0,518 "id": tool_id,519 "type": "function",520 "function": {521 "name": tool_name,522 "arguments": chunk["choices"][0]["text"],523 },524 }525 ],526 },527 }528 ],529 }530 531 if id_ is not None and created is not None and model is not None:532 yield {533 "id": id_,534 "object": "chat.completion.chunk",535 "created": created,536 "model": model,537 "choices": [538 {539 "index": 0,540 "finish_reason": "tool_calls",541 "logprobs": None,542 "delta": {543 "role": None,544 "content": None,545 "function_call": None,546 "tool_calls": None,547 },548 }549 ],550 }551 552 return _stream_response_to_function_stream(chunks)553 554 555def chat_formatter_to_chat_completion_handler(556 chat_formatter: ChatFormatter,557) -> LlamaChatCompletionHandler:558 def chat_completion_handler(559 *,560 llama: llama.Llama,561 messages: List[llama_types.ChatCompletionRequestMessage],562 functions: Optional[List[llama_types.ChatCompletionFunction]] = None,563 function_call: Optional[llama_types.ChatCompletionRequestFunctionCall] = None,564 tools: Optional[List[llama_types.ChatCompletionTool]] = None,565 tool_choice: Optional[llama_types.ChatCompletionToolChoiceOption] = None,566 temperature: float = 0.2,567 top_p: float = 0.95,568 top_k: int = 40,569 min_p: float = 0.05,570 typical_p: float = 1.0,571 stream: bool = False,572 stop: Optional[Union[str, List[str]]] = [],573 seed: Optional[int] = None,574 response_format: Optional[575 llama_types.ChatCompletionRequestResponseFormat576 ] = None,577 max_tokens: Optional[int] = None,578 presence_penalty: float = 0.0,579 frequency_penalty: float = 0.0,580 repeat_penalty: float = 1.1,581 tfs_z: float = 1.0,582 mirostat_mode: int = 0,583 mirostat_tau: float = 5.0,584 mirostat_eta: float = 0.1,585 model: Optional[str] = None,586 logits_processor: Optional[llama.LogitsProcessorList] = None,587 grammar: Optional[llama.LlamaGrammar] = None,588 logit_bias: Optional[Dict[str, float]] = None,589 logprobs: Optional[bool] = None,590 top_logprobs: Optional[int] = None,591 **kwargs, # type: ignore592 ) -> Union[593 llama_types.CreateChatCompletionResponse,594 Iterator[llama_types.CreateChatCompletionStreamResponse],595 ]:596 result = chat_formatter(597 messages=messages,598 functions=functions,599 function_call=function_call,600 tools=tools,601 tool_choice=tool_choice,602 )603 prompt = llama.tokenize(604 result.prompt.encode("utf-8"),605 add_bos=not result.added_special,606 special=True,607 )608 if result.stop is not None:609 stop = [] if stop is None else [stop] if isinstance(stop, str) else stop610 rstop = result.stop if isinstance(result.stop, list) else [result.stop]611 stop = stop + rstop612 613 stopping_criteria = None614 if result.stopping_criteria is not None:615 stopping_criteria = result.stopping_criteria616 617 if response_format is not None and response_format["type"] == "json_object":618 grammar = _grammar_for_response_format(619 response_format, verbose=llama.verbose620 )621 622 # Convert legacy functions to tools623 if functions is not None:624 tools = [625 {626 "type": "function",627 "function": function,628 }629 for function in functions630 ]631 632 # Convert legacy function_call to tool_choice633 if function_call is not None:634 if isinstance(function_call, str) and (635 function_call == "none" or function_call == "auto"636 ):637 tool_choice = function_call638 if isinstance(function_call, dict) and "name" in function_call:639 tool_choice = {640 "type": "function",641 "function": {642 "name": function_call["name"],643 },644 }645 646 tool = None647 if (648 tool_choice is not None649 and isinstance(tool_choice, dict)650 and tools is not None651 ):652 name = tool_choice["function"]["name"]653 tool = next((t for t in tools if t["function"]["name"] == name), None)654 if tool is None:655 raise ValueError(f"Tool choice '{name}' not found in tools.")656 schema = tool["function"]["parameters"]657 try:658 # create grammar from json schema659 grammar = llama_grammar.LlamaGrammar.from_json_schema(660 json.dumps(schema), verbose=llama.verbose661 )662 except Exception as e:663 if llama.verbose:664 print(str(e), file=sys.stderr)665 grammar = llama_grammar.LlamaGrammar.from_string(666 llama_grammar.JSON_GBNF, verbose=llama.verbose667 )668 669 completion_or_chunks = llama.create_completion(670 prompt=prompt,671 temperature=temperature,672 top_p=top_p,673 top_k=top_k,674 min_p=min_p,675 typical_p=typical_p,676 logprobs=top_logprobs if logprobs else None,677 stream=stream,678 stop=stop,679 seed=seed,680 max_tokens=max_tokens,681 presence_penalty=presence_penalty,682 frequency_penalty=frequency_penalty,683 repeat_penalty=repeat_penalty,684 tfs_z=tfs_z,685 mirostat_mode=mirostat_mode,686 mirostat_tau=mirostat_tau,687 mirostat_eta=mirostat_eta,688 model=model,689 logits_processor=logits_processor,690 stopping_criteria=stopping_criteria,691 grammar=grammar,692 logit_bias=logit_bias,693 )694 if tool is not None:695 tool_name = tool["function"]["name"]696 return _convert_completion_to_chat_function(697 tool_name, completion_or_chunks, stream698 )699 return _convert_completion_to_chat(completion_or_chunks, stream=stream)700 701 return chat_completion_handler702 703 704def hf_autotokenizer_to_chat_formatter(705 pretrained_model_name_or_path: Union[str, os.PathLike[str]]706) -> ChatFormatter:707 # https://huggingface.co/docs/transformers/main/chat_templating708 # https://huggingface.co/mistralai/Mistral-7B-Instruct-v0.1#instruction-format709 # https://huggingface.co/mistralai/Mistral-7B-Instruct-v0.1/blob/main/tokenizer_config.json710 from transformers import AutoTokenizer # type: ignore711 712 tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path) # type: ignore713 714 def format_autotokenizer(715 messages: List[llama_types.ChatCompletionRequestMessage],716 **kwargs: Any,717 ) -> ChatFormatterResponse:718 tokenizer.use_default_system_prompt = False # type: ignore719 prompt: str = tokenizer.apply_chat_template(messages, tokenize=False) # type: ignore720 assert isinstance(prompt, str)721 # Return formatted prompt and eos token by default722 return ChatFormatterResponse(723 prompt=prompt, stop=tokenizer.eos_token, added_special=True724 )725 726 return format_autotokenizer727 728 729def hf_autotokenizer_to_chat_completion_handler(730 pretrained_model_name_or_path: Union[str, os.PathLike[str]]731) -> LlamaChatCompletionHandler:732 chat_formatter = hf_autotokenizer_to_chat_formatter(pretrained_model_name_or_path)733 return chat_formatter_to_chat_completion_handler(chat_formatter)734 735 736def hf_tokenizer_config_to_chat_formatter(737 tokenizer_config: Dict[str, Any],738 add_generation_prompt: bool = True,739) -> ChatFormatter:740 assert isinstance(tokenizer_config, dict)741 742 assert "chat_template" in tokenizer_config743 assert isinstance(tokenizer_config["chat_template"], str)744 chat_template = tokenizer_config["chat_template"]745 746 assert "bos_token" in tokenizer_config747 assert isinstance(tokenizer_config["bos_token"], str)748 bos_token = tokenizer_config["bos_token"]749 750 assert "eos_token" in tokenizer_config751 assert isinstance(tokenizer_config["eos_token"], str)752 eos_token = tokenizer_config["eos_token"]753 754 env = ImmutableSandboxedEnvironment(755 trim_blocks=True,756 lstrip_blocks=True,757 ).from_string(chat_template)758 759 def format_tokenizer_config(760 messages: List[llama_types.ChatCompletionRequestMessage],761 **kwargs: Any,762 ) -> ChatFormatterResponse:763 # TODO: veryify this is correct764 # Add a blank assistant message to the end of the messages to prompt the model to generate a response765 if add_generation_prompt:766 messages = [767 *messages,768 llama_types.ChatCompletionRequestAssistantMessage(769 role="assistant", content=""770 ),771 ]772 prompt = env.render(773 messages=messages,774 bos_token=bos_token,775 eos_token=eos_token,776 )777 return ChatFormatterResponse(778 prompt=prompt, stop=[eos_token, bos_token], added_special=True779 )780 781 return format_tokenizer_config782 783 784def hf_tokenizer_config_to_chat_completion_handler(785 tokenizer_config: Dict[str, Any],786 add_generation_prompt: bool = True,787) -> LlamaChatCompletionHandler:788 chat_formatter = hf_tokenizer_config_to_chat_formatter(789 tokenizer_config, add_generation_prompt=add_generation_prompt790 )791 return chat_formatter_to_chat_completion_handler(chat_formatter)792 793 794def guess_chat_format_from_gguf_metadata(metadata: Dict[str, str]) -> Optional[str]:795 if "tokenizer.chat_template" not in metadata:796 return None797 798 if metadata["tokenizer.chat_template"] == CHATML_CHAT_TEMPLATE:799 return "chatml"800 801 if (802 metadata["tokenizer.chat_template"] == MISTRAL_INSTRUCT_CHAT_TEMPLATE803 or metadata["tokenizer.chat_template"] == MIXTRAL_INSTRUCT_CHAT_TEMPLATE804 ):805 return "mistral-instruct"806 807 if metadata["tokenizer.chat_template"] == LLAMA3_INSTRUCT_CHAT_TEMPLATE:808 return "llama-3"809 810 return None811 812 813### Utility functions for formatting chat prompts ###814# TODO: Replace these with jinja2 templates815 816 817def _get_system_message(818 messages: List[llama_types.ChatCompletionRequestMessage],819) -> str:820 """Get the first system message."""821 for message in messages:822 if message["role"] == "system":823 return message["content"] or ""824 return ""825 826 827def _map_roles(828 messages: List[llama_types.ChatCompletionRequestMessage],829 role_map: Dict[str, str],830) -> List[Tuple[str, Optional[str]]]:831 """Map the message roles."""832 output: List[Tuple[str, Optional[str]]] = []833 for message in messages:834 role = message["role"]835 if role in role_map:836 content: str | None = (837 message["content"] if isinstance(message["content"], str) else None838 )839 output.append((role_map[role], content))840 return output841 842 843def _format_llama2(844 system_message: str, messages: List[Tuple[str, Optional[str]]], sep: str, sep2: str845) -> str:846 """Format the prompt with the llama2 style."""847 seps = [sep, sep2]848 ret = system_message + sep849 for i, (role, message) in enumerate(messages):850 if system_message and i == 0:851 m = message or ""852 ret += m + seps[i % 2]853 elif message:854 ret += role + message + " " + seps[i % 2]855 else:856 ret += role + " "857 return ret858 859 860def _format_add_colon_single(861 system_message: str, messages: List[Tuple[str, Optional[str]]], sep: str862) -> str:863 """Format the prompt with the add-colon-single style."""864 ret = system_message + sep865 for role, message in messages:866 if message:867 ret += role + ": " + message + sep868 else:869 ret += role + ":"870 return ret871 872 873def _format_add_colon_two(874 system_message: str, messages: List[Tuple[str, Optional[str]]], sep: str, sep2: str875) -> str:876 """Format the prompt with the add-colon-two style."""877 seps = [sep, sep2]878 ret = system_message + seps[0]879 for i, (role, message) in enumerate(messages):880 if message:881 ret += role + ": " + message + seps[i % 2]882 else:883 ret += role + ":"884 return ret885 886 887def _format_no_colon_single(888 system_message: str, messages: List[Tuple[str, Optional[str]]], sep: str889) -> str:890 """Format the prompt with the no-colon-single style."""891 ret = system_message892 for role, message in messages:893 if message:894 ret += role + message + sep895 else:896 ret += role897 return ret898 899 900def _format_add_colon_space_single(901 system_message: str, messages: List[Tuple[str, Optional[str]]], sep: str902) -> str:903 """Format the prompt with the add-colon-space-single style."""904 ret = system_message + sep905 for role, message in messages:906 if message:907 ret += role + ": " + message + sep908 else:909 ret += role + ": " # must be end with a space910 return ret911 912 913def _format_chatml(914 system_message: str, messages: List[Tuple[str, Optional[str]]], sep: str915) -> str:916 """Format the prompt with the chatml style."""917 ret = "" if system_message == "" else system_message + sep + "\n"918 for role, message in messages:919 if message:920 ret += role + "\n" + message + sep + "\n"921 else:922 ret += role + "\n"923 return ret924 925 926def _format_chatglm3(927 system_message: str, messages: List[Tuple[str, Optional[str]]], sep: str928) -> str:929 """Format the prompt with the chatglm3 style."""930 ret = ""931 if system_message:932 ret += system_message933 for role, message in messages:934 if message:935 ret += role + "\n" + " " + message936 else:937 ret += role938 return ret939 940 941def _grammar_for_json(verbose: bool = False):942 return llama_grammar.LlamaGrammar.from_string(943 llama_grammar.JSON_GBNF, verbose=verbose944 )945 946 947def _grammar_for_json_schema(948 schema: str, verbose: bool = False, fallback_to_json: bool = True949):950 try:951 return llama_grammar.LlamaGrammar.from_json_schema(schema, verbose=verbose)952 except Exception as e:953 if fallback_to_json:954 return _grammar_for_json(verbose=verbose)955 else:956 raise e957 958 959def _grammar_for_response_format(960 response_format: llama_types.ChatCompletionRequestResponseFormat,961 verbose: bool = False,962):963 if response_format["type"] != "json_object":964 return None965 966 if "schema" in response_format:967 return _grammar_for_json_schema(968 json.dumps(response_format["schema"]), verbose=verbose969 )970 else:971 return _grammar_for_json(verbose=verbose)972 973 974### Chat Formats ###975 976 977def register_chat_format(name: str):978 def decorator(f: ChatFormatter):979 chat_completion_handler = chat_formatter_to_chat_completion_handler(f)980 LlamaChatCompletionHandlerRegistry().register_chat_completion_handler(981 name, chat_completion_handler982 )983 return f984 985 return decorator986 987 988# see https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/tokenization_llama.py989# system prompt is "embedded" in the first message990@register_chat_format("llama-2")991def format_llama2(992 messages: List[llama_types.ChatCompletionRequestMessage],993 **kwargs: Any,994) -> ChatFormatterResponse:995 _system_template = "[INST] <<SYS>>\n{system_message}\n<</SYS>>"996 _roles = dict(user="<s>[INST]", assistant="[/INST]")997 _messages = _map_roles(messages, _roles)998 system_message = _get_system_message(messages)999 if system_message:1000 system_message = _system_template.format(system_message=system_message)1001 _prompt = _format_llama2(system_message, _messages, " ", "</s>") + "[/INST]"1002 return ChatFormatterResponse(prompt=_prompt)1003 1004 1005# Chat format for Llama-3 models, see more details at:1006# https://github.com/meta-llama/llama3/blob/main/llama/tokenizer.py#L202-L2291007@register_chat_format("llama-3")1008def format_llama3(1009 messages: List[llama_types.ChatCompletionRequestMessage],1010 **kwargs: Any,1011) -> ChatFormatterResponse:1012 _roles = dict(1013 system="<|start_header_id|>system<|end_header_id|>\n\n",1014 user="<|start_header_id|>user<|end_header_id|>\n\n",1015 assistant="<|start_header_id|>assistant<|end_header_id|>\n\n",1016 )1017 _sep = "<|eot_id|>"1018 _messages = _map_roles(messages, _roles)1019 _messages.append((_roles["assistant"], None))1020 _prompt = _format_no_colon_single("", _messages, _sep)1021 return ChatFormatterResponse(prompt=_prompt, stop=_sep)1022 1023 1024@register_chat_format("alpaca")1025def format_alpaca(1026 messages: List[llama_types.ChatCompletionRequestMessage],1027 **kwargs: Any,1028) -> ChatFormatterResponse:1029 _roles = dict(user="### Instruction", assistant="### Response")1030 _sep = "\n\n"1031 _sep2 = "</s>"1032 system_message = _get_system_message(messages)1033 _messages = _map_roles(messages, _roles)1034 _prompt = _format_add_colon_two(system_message, _messages, _sep, _sep2)1035 return ChatFormatterResponse(prompt=_prompt)1036 1037 1038@register_chat_format("qwen")1039def format_qwen(1040 messages: List[llama_types.ChatCompletionRequestMessage],1041 **kwargs: Any,1042) -> ChatFormatterResponse:1043 _roles = dict(user="<|im_start|>user", assistant="<|im_start|>assistant")1044 system_message = _get_system_message(messages) or "You are a helpful assistant."1045 system_template = "<|im_start|>system\n{system_message}"1046 system_message = system_template.format(system_message=system_message)1047 _messages = _map_roles(messages, _roles)1048 _messages.append((_roles["assistant"], None))1049 _sep = "<|im_end|>"1050 _prompt = _format_chatml(system_message, _messages, _sep)1051 _sep2 = "<|endoftext|>"1052 return ChatFormatterResponse(prompt=_prompt, stop=_sep2)1053 1054 1055@register_chat_format("vicuna")1056def format(1057 messages: List[llama_types.ChatCompletionRequestMessage],1058 **kwargs: Any,1059) -> ChatFormatterResponse:1060 _system_message = "A chat between a curious user and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the user's questions."1061 _roles = dict(user="USER", assistant="ASSISTANT")1062 _sep = " "1063 _sep2 = "</s>"1064 system_message = _system_message1065 _messages = _map_roles(messages, _roles)1066 _messages.append((_roles["assistant"], None))1067 _prompt = _format_add_colon_two(system_message, _messages, _sep, _sep2)1068 return ChatFormatterResponse(prompt=_prompt)1069 1070 1071@register_chat_format("oasst_llama")1072def format_oasst_llama(1073 messages: List[llama_types.ChatCompletionRequestMessage],1074 **kwargs: Any,1075) -> ChatFormatterResponse:1076 _system_template = "[INST] <<SYS>>\n{system_message}\n<</SYS>>\n\n"1077 _roles = dict(user="<|prompter|>", assistant="<|assistant|>")1078 _sep = "</s>"1079 system_message = _get_system_message(messages)1080 system_message = _system_template.format(system_message=system_message)1081 _messages = _map_roles(messages, _roles)1082 _messages.append((_roles["assistant"], None))1083 _prompt = _format_no_colon_single(system_message, _messages, _sep)1084 return ChatFormatterResponse(prompt=_prompt)1085 1086 1087@register_chat_format("baichuan-2")1088def format_baichuan2(1089 messages: List[llama_types.ChatCompletionRequestMessage],1090 **kwargs: Any,1091) -> ChatFormatterResponse:1092 _system_template = "{system_message}"1093 _roles = dict(user="<reserved_106>", assistant="<reserved_107>")1094 _sep = ""1095 system_message = _get_system_message(messages)1096 system_message = _system_template.format(system_message=system_message)1097 _messages = _map_roles(messages, _roles)1098 _messages.append((_roles["assistant"], None))1099 _prompt = _format_no_colon_single(system_message, _messages, _sep)1100 return ChatFormatterResponse(prompt=_prompt)1101 1102 1103@register_chat_format("baichuan")1104def format_baichuan(1105 messages: List[llama_types.ChatCompletionRequestMessage],1106 **kwargs: Any,1107) -> ChatFormatterResponse:1108 _system_template = "{system_message}"1109 _roles = dict(user="<reserved_102>", assistant="<reserved_103>")1110 _sep = ""1111 system_message = _get_system_message(messages)1112 system_message = _system_template.format(system_message=system_message)1113 _messages = _map_roles(messages, _roles)1114 _messages.append((_roles["assistant"], None))1115 _prompt = _format_no_colon_single(system_message, _messages, _sep)1116 return ChatFormatterResponse(prompt=_prompt)1117 1118 1119@register_chat_format("openbuddy")1120def format_openbuddy(1121 messages: List[llama_types.ChatCompletionRequestMessage],1122 **kwargs: Any,1123) -> ChatFormatterResponse:1124 _system_message = """You are a helpful, respectful and honest INTP-T AI Assistant named Buddy. You are talking to a human User.1125Always answer as helpfully and logically as possible, while being safe. Your answers should not include any harmful, political, religious, unethical, racist, sexist, toxic, dangerous, or illegal content. Please ensure that your responses are socially unbiased and positive in nature.1126If a question does not make any sense, or is not factually coherent, explain why instead of answering something not correct. If you don't know the answer to a question, please don't share false information.1127You can speak fluently in many languages, for example: English, Chinese.1128You cannot access the internet, but you have vast knowledge, cutoff: 2021-09.1129You are trained by OpenBuddy team, (https://openbuddy.ai, https://github.com/OpenBuddy/OpenBuddy), you are based on LLaMA and Falcon transformers model, not related to GPT or OpenAI.1130 1131"""1132 _roles = dict(user="User", assistant="Assistant")1133 _sep = "\n"1134 system_message = _system_message1135 _messages = _map_roles(messages, _roles)1136 _messages.append((_roles["assistant"], None))1137 _prompt = _format_add_colon_single(system_message, _messages, _sep)1138 return ChatFormatterResponse(prompt=_prompt)1139 1140 1141@register_chat_format("redpajama-incite")1142def format_redpajama_incite(1143 messages: List[llama_types.ChatCompletionRequestMessage],1144 **kwargs: Any,1145) -> ChatFormatterResponse:1146 _system_message = _get_system_message(messages)1147 _roles = dict(user="<human>", assistant="<bot>")1148 _sep = "\n"1149 _stop = "<human>"1150 system_message = _system_message1151 _messages = _map_roles(messages, _roles)1152 _messages.append((_roles["assistant"], None))1153 _prompt = _format_add_colon_single(system_message, _messages, _sep)1154 return ChatFormatterResponse(prompt=_prompt, stop=_stop)1155 1156 1157@register_chat_format("snoozy")1158def format_snoozy(1159 messages: List[llama_types.ChatCompletionRequestMessage],1160 **kwargs: Any,1161) -> ChatFormatterResponse:1162 system_template = "### Instruction:\n{system_message}"1163 default_system_message = "The prompt below is a question to answer, a task to complete, or a conversation to respond to; decide which and write an appropriate response."1164 _system_message = _get_system_message(messages)1165 _system_message = (1166 _system_message if _system_message != "" else default_system_message1167 )1168 system_message = system_template.format(system_message=_system_message)1169 _roles = dict(user="### Prompt", assistant="### Response")1170 _sep = "\n"1171 _stop = "###"1172 system_message = _system_message1173 _messages = _map_roles(messages, _roles)1174 _messages.append((_roles["assistant"], None))1175 _prompt = _format_add_colon_single(system_message, _messages, _sep)1176 return ChatFormatterResponse(prompt=_prompt, stop=_stop)1177 1178 1179@register_chat_format("phind")1180def format_phind(1181 messages: List[llama_types.ChatCompletionRequestMessage],1182 **kwargs: Any,1183) -> ChatFormatterResponse:1184 _roles = dict(user="### User Message", assistant="### Assistant")1185 _sep = "\n\n"1186 _system_message = "### System Prompt\nYou are an intelligent programming assistant."1187 _messages = _map_roles(messages, _roles)1188 _messages.append((_roles["assistant"], None))1189 _prompt = _format_add_colon_single(_system_message, _messages, _sep)1190 return ChatFormatterResponse(prompt=_prompt)1191 1192 1193@register_chat_format("intel")1194def format_intel(1195 messages: List[llama_types.ChatCompletionRequestMessage],1196 **kwargs: Any,1197) -> ChatFormatterResponse:1198 _roles = dict(user="### User:", assistant="### Assistant:")1199 _sep = "\n"1200 _system_message = "### System:\n{system_message}"