bnitokyo/InternVL
0
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 