CoolFace
Apppublic

q-future/Co-Instruct

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
29likes
conversation.py301 linesDownload Raw Back to mplug_owl2
1import dataclasses2from enum import auto, Enum3from typing import List, Tuple4from mplug_owl2.constants import DEFAULT_IMAGE_TOKEN5 6class SeparatorStyle(Enum):7    """Different separator style."""8    SINGLE = auto()9    TWO = auto()10    TWO_NO_SYS = auto()11    MPT = auto()12    PLAIN = auto()13    LLAMA_2 = auto()14 15 16@dataclasses.dataclass17class Conversation:18    """A class that keeps all conversation history."""19    system: str20    roles: List[str]21    messages: List[List[str]]22    offset: int23    sep_style: SeparatorStyle = SeparatorStyle.SINGLE24    sep: str = "###"25    sep2: str = None26    version: str = "Unknown"27 28    skip_next: bool = False29 30    def get_prompt(self):31        messages = self.messages32        if len(messages) > 0 and type(messages[0][1]) is tuple:33            messages = self.messages.copy()34            init_role, init_msg = messages[0].copy()35            # init_msg = init_msg[0].replace("<image>", "").strip()36            # if 'mmtag' in self.version:37            #     messages[0] = (init_role, init_msg)38            #     messages.insert(0, (self.roles[0], "<Image><image></Image>"))39            #     messages.insert(1, (self.roles[1], "Received."))40            # else:41            #     messages[0] = (init_role, "<image>\n" + init_msg)42            init_msg = init_msg[0].replace(DEFAULT_IMAGE_TOKEN, "").strip()43            messages[0] = (init_role, DEFAULT_IMAGE_TOKEN + init_msg)44 45        if self.sep_style == SeparatorStyle.SINGLE:46            ret = self.system + self.sep47            for role, message in messages:48                if message:49                    if type(message) is tuple:50                        message, _, _ = message51                    ret += role + ": " + message + self.sep52                else:53                    ret += role + ":"54        elif self.sep_style == SeparatorStyle.TWO:55            seps = [self.sep, self.sep2]56            ret = self.system + seps[0]57            for i, (role, message) in enumerate(messages):58                if message:59                    if type(message) is tuple:60                        message, _, _ = message61                    ret += role + ": " + message + seps[i % 2]62                else:63                    ret += role + ":"64        elif self.sep_style == SeparatorStyle.TWO_NO_SYS:65            seps = [self.sep, self.sep2]66            ret = ""67            for i, (role, message) in enumerate(messages):68                if message:69                    if type(message) is tuple:70                        message, _, _ = message71                    ret += role + ": " + message + seps[i % 2]72                else:73                    ret += role + ":"74        elif self.sep_style == SeparatorStyle.MPT:75            ret = self.system + self.sep76            for role, message in messages:77                if message:78                    if type(message) is tuple:79                        message, _, _ = message80                    ret += role + message + self.sep81                else:82                    ret += role83        elif self.sep_style == SeparatorStyle.LLAMA_2:84            wrap_sys = lambda msg: f"<<SYS>>\n{msg}\n<</SYS>>\n\n"85            wrap_inst = lambda msg: f"[INST] {msg} [/INST]"86            ret = ""87 88            for i, (role, message) in enumerate(messages):89                if i == 0:90                    assert message, "first message should not be none"91                    assert role == self.roles[0], "first message should come from user"92                if message:93                    if type(message) is tuple:94                        message, _, _ = message95                    if i == 0: message = wrap_sys(self.system) + message96                    if i % 2 == 0:97                        message = wrap_inst(message)98                        ret += self.sep + message99                    else:100                        ret += " " + message + " " + self.sep2101                else:102                    ret += ""103            ret = ret.lstrip(self.sep)104        elif self.sep_style == SeparatorStyle.PLAIN:105            seps = [self.sep, self.sep2]106            ret = self.system107            for i, (role, message) in enumerate(messages):108                if message:109                    if type(message) is tuple:110                        message, _, _ = message111                    ret += message + seps[i % 2]112                else:113                    ret += ""114        else:115            raise ValueError(f"Invalid style: {self.sep_style}")116 117        return ret118 119    def append_message(self, role, message):120        self.messages.append([role, message])121 122    def get_images(self, return_pil=False):123        images = []124        for i, (role, msg) in enumerate(self.messages[self.offset:]):125            if i % 2 == 0:126                if type(msg) is tuple:127                    import base64128                    from io import BytesIO129                    from PIL import Image130                    msg, image, image_process_mode = msg131                    if image_process_mode == "Pad":132                        def expand2square(pil_img, background_color=(122, 116, 104)):133                            width, height = pil_img.size134                            if width == height:135                                return pil_img136                            elif width > height:137                                result = Image.new(pil_img.mode, (width, width), background_color)138                                result.paste(pil_img, (0, (width - height) // 2))139                                return result140                            else:141                                result = Image.new(pil_img.mode, (height, height), background_color)142                                result.paste(pil_img, ((height - width) // 2, 0))143                                return result144                        image = expand2square(image)145                    elif image_process_mode in ["Default", "Crop"]:146                        pass147                    elif image_process_mode == "Resize":148                        image = image.resize((336, 336))149                    else:150                        raise ValueError(f"Invalid image_process_mode: {image_process_mode}")151                    max_hw, min_hw = max(image.size), min(image.size)152                    aspect_ratio = max_hw / min_hw153                    max_len, min_len = 800, 400154                    shortest_edge = int(min(max_len / aspect_ratio, min_len, min_hw))155                    longest_edge = int(shortest_edge * aspect_ratio)156                    W, H = image.size157                    if longest_edge != max(image.size):158                        if H > W:159                            H, W = longest_edge, shortest_edge160                        else:161                            H, W = shortest_edge, longest_edge162                        image = image.resize((W, H))163                    if return_pil:164                        images.append(image)165                    else:166                        buffered = BytesIO()167                        image.save(buffered, format="PNG")168                        img_b64_str = base64.b64encode(buffered.getvalue()).decode()169                        images.append(img_b64_str)170        return images171 172    def to_gradio_chatbot(self):173        ret = []174        for i, (role, msg) in enumerate(self.messages[self.offset:]):175            if i % 2 == 0:176                if type(msg) is tuple:177                    import base64178                    from io import BytesIO179                    msg, image, image_process_mode = msg180                    max_hw, min_hw = max(image.size), min(image.size)181                    aspect_ratio = max_hw / min_hw182                    max_len, min_len = 800, 400183                    shortest_edge = int(min(max_len / aspect_ratio, min_len, min_hw))184                    longest_edge = int(shortest_edge * aspect_ratio)185                    W, H = image.size186                    if H > W:187                        H, W = longest_edge, shortest_edge188                    else:189                        H, W = shortest_edge, longest_edge190                    image = image.resize((W, H))191                    buffered = BytesIO()192                    image.save(buffered, format="JPEG")193                    img_b64_str = base64.b64encode(buffered.getvalue()).decode()194                    img_str = f'<img src="data:image/png;base64,{img_b64_str}" alt="user upload image" />'195                    msg = img_str + msg.replace('<|image|>', '').strip()196                    ret.append([msg, None])197                else:198                    ret.append([msg, None])199            else:200                ret[-1][-1] = msg201        return ret202 203    def copy(self):204        return Conversation(205            system=self.system,206            roles=self.roles,207            messages=[[x, y] for x, y in self.messages],208            offset=self.offset,209            sep_style=self.sep_style,210            sep=self.sep,211            sep2=self.sep2,212            version=self.version)213 214    def dict(self):215        if len(self.get_images()) > 0:216            return {217                "system": self.system,218                "roles": self.roles,219                "messages": [[x, y[0] if type(y) is tuple else y] for x, y in self.messages],220                "offset": self.offset,221                "sep": self.sep,222                "sep2": self.sep2,223            }224        return {225            "system": self.system,226            "roles": self.roles,227            "messages": self.messages,228            "offset": self.offset,229            "sep": self.sep,230            "sep2": self.sep2,231        }232 233 234conv_vicuna_v0 = Conversation(235    system="A chat between a curious human and an artificial intelligence assistant. "236           "The assistant gives helpful, detailed, and polite answers to the human's questions.",237    roles=("Human", "Assistant"),238    messages=(239        ("Human", "What are the key differences between renewable and non-renewable energy sources?"),240        ("Assistant",241            "Renewable energy sources are those that can be replenished naturally in a relatively "242            "short amount of time, such as solar, wind, hydro, geothermal, and biomass. "243            "Non-renewable energy sources, on the other hand, are finite and will eventually be "244            "depleted, such as coal, oil, and natural gas. Here are some key differences between "245            "renewable and non-renewable energy sources:\n"246            "1. Availability: Renewable energy sources are virtually inexhaustible, while non-renewable "247            "energy sources are finite and will eventually run out.\n"248            "2. Environmental impact: Renewable energy sources have a much lower environmental impact "249            "than non-renewable sources, which can lead to air and water pollution, greenhouse gas emissions, "250            "and other negative effects.\n"251            "3. Cost: Renewable energy sources can be more expensive to initially set up, but they typically "252            "have lower operational costs than non-renewable sources.\n"253            "4. Reliability: Renewable energy sources are often more reliable and can be used in more remote "254            "locations than non-renewable sources.\n"255            "5. Flexibility: Renewable energy sources are often more flexible and can be adapted to different "256            "situations and needs, while non-renewable sources are more rigid and inflexible.\n"257            "6. Sustainability: Renewable energy sources are more sustainable over the long term, while "258            "non-renewable sources are not, and their depletion can lead to economic and social instability.\n")259    ),260    offset=2,261    sep_style=SeparatorStyle.SINGLE,262    sep="###",263)264 265conv_vicuna_v1 = Conversation(266    system="A chat between a curious user and an artificial intelligence assistant. "267    "The assistant gives helpful, detailed, and polite answers to the user's questions.",268    roles=("USER", "ASSISTANT"),269    version="v1",270    messages=(),271    offset=0,272    sep_style=SeparatorStyle.TWO,273    sep=" ",274    sep2="</s>",275)276 277conv_mplug_owl2 = Conversation(278    system="A chat between a curious human and an artificial intelligence assistant. "279           "The assistant gives helpful, detailed, and polite answers to the human's questions.",280    roles=("USER", "ASSISTANT"),281    version="v1",282    messages=(),283    offset=0,284    sep_style=SeparatorStyle.TWO_NO_SYS,285    sep=" ",286    sep2="</s>",287)288 289# default_conversation = conv_vicuna_v1290default_conversation = conv_mplug_owl2291conv_templates = {292    "default": conv_vicuna_v0,293    "v0": conv_vicuna_v0,294    "v1": conv_vicuna_v1,295    "vicuna_v1": conv_vicuna_v1,296    "mplug_owl2": conv_mplug_owl2,297}298 299 300if __name__ == "__main__":301    print(default_conversation.get_prompt())