CoolFace
Modelpublic

GeneZC/MiniChat-2-3B

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
24likes657downloads
conversation.py224 linesDownload Raw Back to root
1"""2Conversation prompt templates.3"""4 5import dataclasses6from enum import auto, Enum7from typing import List, Tuple, Any8 9 10class SeparatorStyle(Enum):11    """Different separator style."""12 13    ADD_COLON_SINGLE = auto()14    ADD_COLON_TWO = auto()15    NO_COLON_SINGLE = auto()16    BAIZE = auto()17    PHOENIX = auto()18    MINICHAT = auto()19 20 21@dataclasses.dataclass22class Conversation:23    """A class that keeps all conversation history."""24 25    # System prompts26    system: str27    # Two roles28    roles: List[str]29    # All messages30    messages: List[List[str]]31    # Offset of few shot examples32    offset: int33    # Separator34    sep_style: SeparatorStyle35    sep: str36    sep2: str = None37    # Stop criteria (the default one is EOS token)38    stop_str: str = None39    # Stops generation if meeting any token in this list40    stop_token_ids: List[int] = None41 42    # Used for the state in the gradio servers.43    # TODO(lmzheng): refactor this44    conv_id: Any = None45    skip_next: bool = False46    model_name: str = None47 48    def get_prompt(self):49        if self.sep_style == SeparatorStyle.ADD_COLON_SINGLE:50            ret = self.system + self.sep51            for role, message in self.messages:52                if message:53                    ret += role + ": " + message + self.sep54                else:55                    ret += role + ": "56            return ret57        elif self.sep_style == SeparatorStyle.ADD_COLON_TWO:58            seps = [self.sep, self.sep2]59            ret = self.system + seps[0]60            for i, (role, message) in enumerate(self.messages):61                if message:62                    ret += role + ": " + message + seps[i % 2]63                else:64                    ret += role + ": "65            return ret66        elif self.sep_style == SeparatorStyle.NO_COLON_SINGLE:67            ret = self.system68            for role, message in self.messages:69                if message:70                    ret += role + message + self.sep71                else:72                    ret += role73            return ret74        elif self.sep_style == SeparatorStyle.BAIZE:75            ret = self.system + "\n"76            for role, message in self.messages:77                if message:78                    ret += role + message + "\n"79                else:80                    ret += role81            return ret82        elif self.sep_style == SeparatorStyle.PHOENIX:83            ret = self.system84            for role, message in self.messages:85                if message:86                    ret += role + ": " + "<s>" + message + "</s>"87                else:88                    ret += role + ": " + "<s>"89            return ret90        elif self.sep_style == SeparatorStyle.MINICHAT:91            ret = self.system92            for role, message in self.messages:93                if message:94                    ret += role + " " + message + "</s>"95                else:96                    ret += role # No space is needed.97            return ret98        else:99            raise ValueError(f"Invalid style: {self.sep_style}")100 101    def append_message(self, role, message):102        self.messages.append([role, message])103 104    def to_gradio_chatbot(self):105        ret = []106        for i, (role, msg) in enumerate(self.messages[self.offset:]):107            if i % 2 == 0:108                ret.append([msg, None])109            else:110                ret[-1][-1] = msg111        return ret112 113    def to_openai_api_messages(self):114        ret = [{"role": "system", "content": self.system}]115 116        for i, (_, msg) in enumerate(self.messages[self.offset:]):117            if i % 2 == 0:118                ret.append({"role": "user", "content": msg})119            else:120                if msg is not None:121                    ret.append({"role": "assistant", "content": msg})122        return ret123 124    def copy(self):125        return Conversation(126            system=self.system,127            roles=self.roles,128            messages=[[x, y] for x, y in self.messages],129            offset=self.offset,130            sep_style=self.sep_style,131            sep=self.sep,132            sep2=self.sep2,133            stop_str=self.stop_str,134            stop_token_ids=self.stop_token_ids,135            conv_id=self.conv_id,136            model_name=self.model_name,137        )138 139    def dict(self):140        return {141            "system": self.system,142            "roles": self.roles,143            "messages": self.messages,144            "offset": self.offset,145            "conv_id": self.conv_id,146            "model_name": self.model_name,147        }148 149 150conv_vicuna = Conversation(151    system="A chat between a curious user and an artificial intelligence assistant. "152    "The assistant gives helpful, detailed, and polite answers to the user's questions.",153    roles=("USER", "ASSISTANT"),154    messages=(),155    offset=0,156    sep_style=SeparatorStyle.ADD_COLON_TWO,157    sep=" ",158    sep2="</s>",159)160 161conv_baize = Conversation(162    system="The following is a conversation between a human and an AI assistant named Baize (named after a mythical creature in Chinese folklore). Baize is an open-source AI assistant developed by UCSD and Sun Yat-Sen University. The human and the AI assistant take turns chatting. Human statements start with [|Human|] and AI assistant statements start with [|AI|]. The AI assistant always provides responses in as much detail as possible, and in Markdown format. The AI assistant always declines to engage with topics, questions and instructions related to unethical, controversial, or sensitive issues. Complete the transcript in exactly that format.\n",163    roles=("[|Human|]", "[|AI|]"),164    messages=(165        ("[|Human|]", "Hello!"),166        ("[|AI|]", "Hi!"),167    ),168    offset=2,169    sep_style=SeparatorStyle.BAIZE,170    sep="\n",171    stop_str="[|Human|]",172)173 174conv_phoenix = Conversation(175    system="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\n",176    roles=("Human", "Assistant"),177    messages=(),178    offset=0,179    sep_style=SeparatorStyle.PHOENIX,180    sep="</s>",181)182 183conv_chatgpt = Conversation(184    system="You are a helpful assistant.",185    roles=("user", "assistant"),186    messages=(),187    offset=0,188    sep_style=None,189    sep=None,190)191 192conv_minichat = Conversation(193    system="‘MiniChat’是一个由‘Beccurio’开发的AI语言模型。下面是人类和MiniChat之间的一段对话。MiniChat的回复应当尽可能详细,并且以Markdown的形式输出。MiniChat应当拒绝参与违背伦理的讨论。</s>",194    roles=("[|User|]", "[|Assistant|]"),195    messages=(),196    offset=0,197    sep_style=SeparatorStyle.MINICHAT,198    sep="</s>",199)200 201 202conv_templates = {203    "vicuna": conv_vicuna,204    "baize": conv_baize,205    "phoenix": conv_phoenix,206    "chatgpt": conv_chatgpt,207    "minichat": conv_minichat,208}209 210def get_default_conv_template(model_name):211    model_name = model_name.lower()212    try:213        ret = conv_templates[model_name]214        return ret.copy()215    except:216        raise NotImplementedError(f"No support for model {model_name}.")217    218 219if __name__ == "__main__":220    conv = conv_templates["minichat"].copy()221    conv.append_message(conv.roles[0], "Write a Python function that checks if a given number is even or odd.")222    conv.append_message(conv.roles[1], None)223    print([conv.get_prompt()])224