q-future/Co-Instruct
29
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())