CoolFace
Apppublic

bnitokyo/InternVL

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
conversation.py307 linesDownload Raw Back to root
1import os2import dataclasses3import base644import copy5import hashlib6import datetime7from io import BytesIO8from PIL import Image9from typing import Any, List, Dict, Union10from dataclasses import field11 12from utils import LOGDIR13 14 15def pil2base64(img: Image.Image) -> str:16    buffered = BytesIO()17    img.save(buffered, format="PNG")18    return base64.b64encode(buffered.getvalue()).decode()19 20 21def resize_img(img: Image.Image, max_len: int, min_len: int) -> Image.Image:22    max_hw, min_hw = max(img.size), min(img.size)23    aspect_ratio = max_hw / min_hw24    # max_len, min_len = 800, 40025    shortest_edge = int(min(max_len / aspect_ratio, min_len, min_hw))26    longest_edge = int(shortest_edge * aspect_ratio)27    W, H = img.size28    if H > W:29        H, W = longest_edge, shortest_edge30    else:31        H, W = shortest_edge, longest_edge32    return img.resize((W, H))33 34 35@dataclasses.dataclass36class Conversation:37    """A class that keeps all conversation history."""38 39    SYSTEM = "system"40    USER = "user"41    ASSISTANT = "assistant"42 43    roles: List[str] = field(44        default_factory=lambda: [45            Conversation.SYSTEM,46            Conversation.USER,47            Conversation.ASSISTANT,48        ]49    )50    mandatory_system_message = "我是书生·万象,英文名是InternVL,是由上海人工智能实验室、清华大学及多家合作单位联合开发的多模态大语言模型。"51    system_message: str = "请尽可能详细地回答用户的问题。"52    messages: List[Dict[str, Any]] = field(default_factory=lambda: [])53    max_image_limit: int = 454    skip_next: bool = False55    streaming_placeholder: str = "▌"56 57    def get_system_message(self):58        return self.mandatory_system_message + "\n\n" + self.system_message59 60    def set_system_message(self, system_message: str):61        self.system_message = system_message62        return self63 64    def get_prompt(self, inlude_image=False):65        send_messages = [{"role": "system", "content": self.get_system_message()}]66        # send_messages = []67        for message in self.messages:68            if message["role"] == self.USER:69                user_message = {70                    "role": self.USER,71                    "content": message["content"],72                }73                if inlude_image and "image" in message:74                    user_message["image"] = []75                    for image in message["image"]:76                        user_message["image"].append(pil2base64(image))77                send_messages.append(user_message)78            elif message["role"] == self.ASSISTANT:79                send_messages.append(80                    {"role": self.ASSISTANT, "content": message["content"]}81                )82            elif message["role"] == self.SYSTEM:83                send_messages.append(84                    {85                        "role": self.SYSTEM,86                        "content": message["content"],87                    }88                )89            else:90                raise ValueError(f"Invalid role: {message['role']}")91        return send_messages92    93    def get_prompt_v2(self, inlude_image=False, max_dynamic_patch=12):94        send_messages = [95            {96                "role": "system", 97                "content": self.get_system_message(),98            }99        ]100        for message in self.messages:101            if message["role"] == self.USER:102                user_message = {103                    "role": self.USER,104                    "content": message["content"],105                }106                if inlude_image and "image" in message:107                    user_message["image"] = []108                    for image in message["image"]:109                        user_message["image"].append(pil2base64(image))110                    111                    content = [{"type": "text", "text": message["content"]}]112                    for image_base64 in user_message["image"]:113                        content.append({114                            "type": "image_url", 115                            "image_url": {116                                "url": f"data:image/jpeg;base64,{image_base64}",117                                "max_dynamic_patch": max_dynamic_patch118                            }119                        })120                    send_messages.append({'role': self.USER, 'content': content})121                else:122                    send_messages.append(user_message)123            elif message["role"] == self.ASSISTANT:124                send_messages.append(125                    {"role": self.ASSISTANT, "content": message["content"]}126                )127            elif message["role"] == self.SYSTEM:128                send_messages.append(129                    {130                        "role": self.SYSTEM,131                        "content": message["content"],132                    }133                )134            else:135                raise ValueError(f"Invalid role: {message['role']}")136        return send_messages137 138    def append_message(139        self,140        role,141        content,142        image_list=None,143    ):144        self.messages.append(145            {146                "role": role,147                "content": content,148                "image": [] if image_list is None else image_list,149                # "filenames": save_filenames,150            }151        )152 153    def get_images(154        self,155        return_copy=False,156        return_base64=False,157        source: Union[str, None] = None,158    ):159        assert source in [self.USER, self.ASSISTANT, None], f"Invalid source: {soure}"160        images = []161        for i, msg in enumerate(self.messages):162            if source and msg["role"] != source:163                continue164 165            for image in msg.get("image", []):166                # org_image = [i.copy() for i in image]167                if return_copy:168                    image = image.copy()169 170                if return_base64:171                    image = pil2base64(image)172 173                images.append(image)174 175        return images176 177    def to_gradio_chatbot(self):178        ret = []179        for i, msg in enumerate(self.messages):180            if msg["role"] == self.SYSTEM:181                continue182 183            alt_str = (184                "user upload image" if msg["role"] == self.USER else "output image"185            )186            image = msg.get("image", [])187            if not isinstance(image, list):188                images = [image]189            else:190                images = image191 192            img_str_list = []193            for i in range(len(images)):194                image = resize_img(195                    images[i],196                    400,197                    200,198                )199                img_b64_str = pil2base64(image)200                W, H = image.size201                img_str = f'<img src="data:image/png;base64,{img_b64_str}" alt="{alt_str}" style="width: {W}px; max-width:none; max-height:none"></img>'202                # img_str = (203                #     f'<img src="data:image/png;base64,{img_b64_str}" alt="{alt_str}" />'204                # )205                img_str_list.append(img_str)206 207            if msg["role"] == self.USER:208                msg_str = " ".join(img_str_list) + msg["content"]209                ret.append([msg_str, None])210            else:211                msg_str = msg["content"] + " ".join(img_str_list)212                ret[-1][-1] = msg_str213        return ret214 215    def update_message(self, role, content, image=None, idx=-1):216        assert len(self.messages) > 0, "No message in the conversation."217 218        idx = (idx + len(self.messages)) % len(self.messages)219 220        assert (221            self.messages[idx]["role"] == role222        ), f"Role mismatch: {role} vs {self.messages[idx]['role']}"223 224        self.messages[idx]["content"] = content225        if image is not None:226            if image not in self.messages[idx]["image"]:227                self.messages[idx]["image"] = []228            if not isinstance(image, list):229                image = [image]230            self.messages[idx]["image"].extend(image)231 232    def return_last_message(self):233        return self.messages[-1]["content"]234 235    def end_of_current_turn(self):236        assert len(self.messages) > 0, "No message in the conversation."237        assert (238            self.messages[-1]["role"] == self.ASSISTANT239        ), f"It should end with the message from assistant instead of {self.messages[-1]['role']}."240 241        if self.messages[-1]["content"][-1] != self.streaming_placeholder:242            return243 244        self.update_message(self.ASSISTANT, self.messages[-1]["content"][:-1], None)245 246    def copy(self):247        return Conversation(248            mandatory_system_message=self.mandatory_system_message,249            system_message=self.system_message,250            roles=copy.deepcopy(self.roles),251            messages=copy.deepcopy(self.messages),252        )253 254    def dict(self):255        """256        all_images = state.get_images()257        all_image_hash = [hashlib.md5(image.tobytes()).hexdigest() for image in all_images]258        t = datetime.datetime.now()259        for image, hash in zip(all_images, all_image_hash):260            filename = os.path.join(261                LOGDIR, "serve_images", f"{t.year}-{t.month:02d}-{t.day:02d}", f"{hash}.jpg"262            )263            if not os.path.isfile(filename):264                os.makedirs(os.path.dirname(filename), exist_ok=True)265                image.save(filename)266        """267        messages = []268        for message in self.messages:269            images = []270            for image in message.get("image", []):271                filename = self.save_image(image)272                images.append(filename)273 274            messages.append(275                {276                    "role": message["role"],277                    "content": message["content"],278                    "image": images,279                }280            )281            if len(images) == 0:282                messages[-1].pop("image")283 284        return {285            "mandatory_system_message": self.mandatory_system_message,286            "system_message": self.system_message,287            "roles": self.roles,288            "messages": messages,289        }290 291    def save_image(self, image: Image.Image) -> str:292        t = datetime.datetime.now()293        image_hash = hashlib.md5(image.tobytes()).hexdigest()294        filename = os.path.join(295            LOGDIR,296            "serve_images",297            f"{t.year}-{t.month:02d}-{t.day:02d}",298            f"{image_hash}.jpg",299        )300        if not os.path.isfile(filename):301            os.makedirs(os.path.dirname(filename), exist_ok=True)302            image.save(filename)303 304        return filename305 306 307