From ca129ab70c7d949e05b97e7704b68d7b93cbf64d Mon Sep 17 00:00:00 2001 From: Dan McPherson Date: Tue, 13 Aug 2024 13:51:07 -0400 Subject: [PATCH] Remove fastchat dependency There were only a couple of remaining calls to fastchat. Models we don't interact with have been removed from support. Signed-off-by: Dan McPherson --- requirements.txt | 1 - src/instructlab/eval/mt_bench_answers.py | 5 +- src/instructlab/eval/mt_bench_common.py | 12 +- src/instructlab/eval/mt_bench_conversation.py | 213 ++++++++++++++++++ .../eval/mt_bench_model_adapter.py | 160 +++++++++++++ 5 files changed, 380 insertions(+), 11 deletions(-) create mode 100644 src/instructlab/eval/mt_bench_conversation.py create mode 100644 src/instructlab/eval/mt_bench_model_adapter.py diff --git a/requirements.txt b/requirements.txt index dd9d302d..9be7cbd8 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -FastChat GitPython>=3.1.42,<4.0.0 shortuuid openai>=1.13.3,<2.0.0 diff --git a/src/instructlab/eval/mt_bench_answers.py b/src/instructlab/eval/mt_bench_answers.py index e80f9962..ac6b98ba 100644 --- a/src/instructlab/eval/mt_bench_answers.py +++ b/src/instructlab/eval/mt_bench_answers.py @@ -6,8 +6,6 @@ import time # Third Party -# TODO need to look into this dependency -from fastchat.model.model_adapter import get_conversation_template # type: ignore import shortuuid import tqdm @@ -20,6 +18,7 @@ load_questions, temperature_config, ) +from .mt_bench_model_adapter import get_conversation_template # type: ignore logger = setup_logger(__name__) @@ -61,7 +60,7 @@ def get_answer( choices = [] for i in range(num_choices): - conv = get_conversation_template(model) + conv = get_conversation_template(model, "granite") turns = [] for j in range(len(question["turns"])): diff --git a/src/instructlab/eval/mt_bench_common.py b/src/instructlab/eval/mt_bench_common.py index 248cfacf..f45bf5f6 100644 --- a/src/instructlab/eval/mt_bench_common.py +++ b/src/instructlab/eval/mt_bench_common.py @@ -13,8 +13,6 @@ import time # Third Party -from fastchat import conversation -from fastchat.model.model_adapter import get_conversation_template # type: ignore import openai # First Party @@ -22,6 +20,8 @@ # Local from .logger_config import setup_logger +from .mt_bench_conversation import Conversation +from .mt_bench_model_adapter import get_conversation_template logger = setup_logger(__name__) @@ -158,7 +158,7 @@ def run_judge_single( rating = -1 system_prompt = judge.prompt_template["system_prompt"] - conv = get_conversation_template(model) + conv = get_conversation_template(model, "mixtral") conv.set_system_message(system_prompt) conv.append_message(conv.roles[0], user_prompt) conv.append_message(conv.roles[1], None) @@ -268,9 +268,7 @@ class Message(TypedDict): role: str -def _get_messages( - conv: conversation.Conversation, merge_system_user_message: bool -) -> list[Message]: +def _get_messages(conv: Conversation, merge_system_user_message: bool) -> list[Message]: messages = conv.to_openai_api_messages() if ( (merge_system_user_message or conv.name == "mistral") @@ -285,7 +283,7 @@ def _get_messages( def chat_completion_openai( openai_client, model, - conv: conversation.Conversation, + conv: Conversation, temperature, max_tokens, merge_system_user_message: bool = False, diff --git a/src/instructlab/eval/mt_bench_conversation.py b/src/instructlab/eval/mt_bench_conversation.py new file mode 100644 index 00000000..05585b90 --- /dev/null +++ b/src/instructlab/eval/mt_bench_conversation.py @@ -0,0 +1,213 @@ +# SPDX-License-Identifier: Apache-2.0 +""" +Conversation prompt templates. +""" + +# Standard +from enum import IntEnum, auto +from typing import Dict, List, Tuple, Union +import dataclasses + + +class SeparatorStyle(IntEnum): + """Separator styles.""" + + ADD_COLON_SINGLE = auto() + ADD_COLON_TWO = auto() + ADD_COLON_SPACE_SINGLE = auto() + NO_COLON_SINGLE = auto() + NO_COLON_TWO = auto() + ADD_NEW_LINE_SINGLE = auto() + LLAMA2 = auto() + DEFAULT = auto() + + +@dataclasses.dataclass +class Conversation: + # pylint: disable=too-many-instance-attributes + """A class that manages prompt templates and keeps all conversation history.""" + + # The name of this template + name: str + # The template of the system prompt + system_template: str = "{system_message}" + # The system message + system_message: str = "" + # The names of two roles + roles: Tuple[str, str] = ("USER", "ASSISTANT") + # All messages. Each item is (role, message). + # Each message is either a string or a tuple of (string, List[image_url]). + messages: List[List[str | None]] = dataclasses.field(default_factory=list) + # The number of few shot examples + offset: int = 0 + # The separator style and configurations + sep_style: SeparatorStyle = SeparatorStyle.ADD_COLON_SINGLE + sep: str | None = "\n" + sep2: str | None = None + # Stop criteria (the default one is EOS token) + stop_str: Union[str, List[str]] | None = None + # Stops generation if meeting any token in this list + stop_token_ids: List[int] | None = None + + def set_system_message(self, system_message: str): + """Set the system message.""" + self.system_message = system_message + + def get_system_message(self): + """return the system message.""" + return self.system_message + + def append_message(self, role: str, message: str | None): + """Append a new message.""" + self.messages.append([role, message]) + + def update_last_message(self, message: str): + """Update the last output. + + The last message is typically set to be None when constructing the prompt, + so we need to update it in-place after getting the response from a model. + """ + self.messages[-1][1] = message + + def to_openai_api_messages(self): + """Convert the conversation to OpenAI chat completion format.""" + if self.system_message == "": + ret = [] + else: + ret = [{"role": "system", "content": self.system_message}] + + for i, (_, msg) in enumerate(self.messages[self.offset :]): + if i % 2 == 0: + ret.append({"role": "user", "content": msg}) + else: + if msg is not None: + ret.append({"role": "assistant", "content": msg}) + return ret + + def copy(self): + return Conversation( + name=self.name, + system_template=self.system_template, + system_message=self.system_message, + roles=self.roles, + messages=[[x, y] for x, y in self.messages], + offset=self.offset, + sep_style=self.sep_style, + sep=self.sep, + sep2=self.sep2, + stop_str=self.stop_str, + stop_token_ids=self.stop_token_ids, + ) + + def dict(self): + return { + "template_name": self.name, + "system_message": self.system_message, + "roles": self.roles, + "messages": self.extract_text_from_messages(), + "offset": self.offset, + } + + +# A global registry for all conversation templates +conv_templates: Dict[str, Conversation] = {} + + +def register_conv_template(template: Conversation, override: bool = False): + """Register a new conversation template.""" + if not override: + assert ( + template.name not in conv_templates + ), f"{template.name} has been registered." + + conv_templates[template.name] = template + + +def get_conv_template(name: str) -> Conversation: + """Get a conversation template.""" + return conv_templates[name].copy() + + +# An empty template for raw conversation. +register_conv_template( + Conversation( + name="raw", + system_message="", + roles=("", ""), + sep_style=SeparatorStyle.NO_COLON_SINGLE, + sep="", + ) +) + + +# api-based default template +register_conv_template( + Conversation( + name="api_based_default", + system_message="", + roles=("user", "assistant"), + sep_style=SeparatorStyle.DEFAULT, + sep=None, + ) +) + + +# ChatGPT default template +register_conv_template( + Conversation( + name="chatgpt", + system_message="You are a helpful assistant.", + roles=("user", "assistant"), + sep_style=SeparatorStyle.DEFAULT, + sep=None, + ) +) + +# Mistral template +# source: https://docs.mistral.ai/llm/mistral-instruct-v0.1#chat-template +register_conv_template( + Conversation( + name="mistral", + system_template="[INST] {system_message}\n", + roles=("[INST]", "[/INST]"), + sep_style=SeparatorStyle.LLAMA2, + sep=" ", + sep2="", + ) +) + +register_conv_template( + Conversation( + name="labrador-chat", + system_template="<|system|>\n{system_message}", + system_message="""You are Labrador, an AI language model developed by IBM DMF (Data Model Factory) Alignment Team. You are a cautious assistant. You carefully follow instructions. You are helpful and harmless and you follow ethical guidelines and promote positive behavior. You always respond to greetings (for example, hi, hello, g'day, morning, afternoon, evening, night, what's up, nice to meet you, sup, etc) with "Hello! I am Labrador, created by the IBM DMF Alignment Team. How can I help you today?". Please do not say anything else and do not start a conversation.""", + roles=("<|user|>", "<|assistant|>"), + sep_style=SeparatorStyle.ADD_NEW_LINE_SINGLE, + sep="\n", + stop_str="<|endoftext|>", + ) +) + +register_conv_template( + Conversation( + name="ibm-generic", + system_template="<|system|>\n{system_message}", + system_message="""You are an AI language model developed by IBM Research. You are a cautious assistant. You carefully follow instructions. You are helpful and harmless and you follow ethical guidelines and promote positive behavior.""", + roles=("<|user|>", "<|assistant|>"), + sep_style=SeparatorStyle.ADD_NEW_LINE_SINGLE, + sep="\n", + stop_str="<|endoftext|>", + ) +) + +register_conv_template( + Conversation( + name="granite-chat", + system_template="<|system|>\n{system_message}", + system_message="""You are Granite Chat, an AI language model developed by IBM. You are a cautious assistant. You carefully follow instructions. You are helpful and harmless and you follow ethical guidelines and promote positive behavior.""", + roles=("<|user|>", "<|assistant|>"), + sep_style=SeparatorStyle.ADD_NEW_LINE_SINGLE, + sep="\n", + stop_str="<|endoftext|>", + ) +) diff --git a/src/instructlab/eval/mt_bench_model_adapter.py b/src/instructlab/eval/mt_bench_model_adapter.py new file mode 100644 index 00000000..67050127 --- /dev/null +++ b/src/instructlab/eval/mt_bench_model_adapter.py @@ -0,0 +1,160 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Model adapter registration.""" + +# Standard +from functools import cache +from typing import List +import abc +import os + +# Local +from .logger_config import setup_logger +from .mt_bench_conversation import Conversation, get_conv_template + +OPENAI_MODEL_LIST = ("gpt-4",) + +logger = setup_logger(__name__) + + +class BaseModelAdapter: + """The base and the default model adapter.""" + + @abc.abstractmethod + def match(self, model_path: str) -> bool: + pass + + @abc.abstractmethod + def get_default_conv_template(self, model_path: str) -> Conversation: + pass + + +# A global registry for all model adapters +model_adapters: List[BaseModelAdapter] = [] + + +def register_model_adapter(cls): + """Register a model adapter.""" + model_adapters.append(cls()) + + +@cache +def get_model_adapter(model_path: str, default_adapter_name: str) -> BaseModelAdapter: + """Get a model adapter for a model_path.""" + model_path_basename = os.path.basename(os.path.normpath(model_path)) + + default_adapter = None + + # Try the basename of model_path at first + for adapter in model_adapters: + if adapter.match(model_path_basename): + return adapter + if adapter.match(default_adapter_name) and default_adapter is None: + default_adapter = adapter + + # Then try the full path + for adapter in model_adapters: + if adapter.match(model_path): + return adapter + + if default_adapter is not None: + logger.warning( + "No valid model adapter for %s, defaulting to %s adapter", + model_path, + default_adapter_name, + ) + return default_adapter + raise ValueError(f"No valid model adapter for {model_path}") + + +def get_conversation_template( + model_path: str, default_adapter_name: str +) -> Conversation: + """Get the default conversation template.""" + adapter = get_model_adapter(model_path, default_adapter_name) + return adapter.get_default_conv_template(model_path) + + +class ChatGPTAdapter(BaseModelAdapter): + """The model adapter for ChatGPT""" + + def match(self, model_path: str): + return model_path in OPENAI_MODEL_LIST + + def get_default_conv_template(self, model_path: str) -> Conversation: + if "browsing" in model_path: + return get_conv_template("api_based_default") + return get_conv_template("chatgpt") + + +class MistralAdapter(BaseModelAdapter): + """The model adapter for Mistral AI models""" + + def match(self, model_path: str): + model_path = model_path.lower() + return ( + "mistral" in model_path + or "mixtral" in model_path + or "prometheus" in model_path + ) + + def get_default_conv_template(self, model_path: str) -> Conversation: + return get_conv_template("mistral") + + +class LabradoriteAdapter(BaseModelAdapter): + """The model adapter for ibm/labradorite-13b""" + + def match(self, model_path: str): + return "labradorite" in model_path.lower() + + def get_default_conv_template(self, model_path: str) -> Conversation: + return get_conv_template("labrador-chat") + + +class MerliniteAdapter(BaseModelAdapter): + """The model adapter for ibm/merlinite-7b and instructlab/merlinite-7b-lab""" + + def match(self, model_path: str): + return "merlinite" in model_path.lower() + + def get_default_conv_template(self, model_path: str) -> Conversation: + return get_conv_template("ibm-generic") + + +class GraniteAdapter(BaseModelAdapter): + """The model adapter for instructlab/granite-7b-lab""" + + def match(self, model_path: str): + model_path = model_path.lower() + return ( + "granite" in model_path + and "granite-old" not in model_path + and "granite-chat" not in model_path + and "granite-code" not in model_path + ) + + def get_default_conv_template(self, model_path: str) -> Conversation: + return get_conv_template("ibm-generic") + + +class LabradorAdapter(BaseModelAdapter): + """The model adapter for ibm/labradorite-13b""" + + def match(self, model_path: str): + model_path = model_path.lower() + return ("granite-chat" in model_path) or ( + "labrador" in model_path and "labradorite" not in model_path + ) + + def get_default_conv_template(self, model_path: str) -> Conversation: + return get_conv_template("granite-chat") + + +# Note: the registration order matters. +# The one registered earlier has a higher matching priority. +register_model_adapter(MistralAdapter) +register_model_adapter(LabradoriteAdapter) +register_model_adapter(MerliniteAdapter) +register_model_adapter(GraniteAdapter) +register_model_adapter(LabradorAdapter) +register_model_adapter(ChatGPTAdapter)