| 162 | |
| 163 | |
| 164 | class PromptExpander: |
| 165 | |
| 166 | def __init__(self, model_name, is_vl=False, device=0, **kwargs): |
| 167 | self.model_name = model_name |
| 168 | self.is_vl = is_vl |
| 169 | self.device = device |
| 170 | |
| 171 | def extend_with_img(self, |
| 172 | prompt, |
| 173 | system_prompt, |
| 174 | image=None, |
| 175 | seed=-1, |
| 176 | *args, |
| 177 | **kwargs): |
| 178 | pass |
| 179 | |
| 180 | def extend(self, prompt, system_prompt, seed=-1, *args, **kwargs): |
| 181 | pass |
| 182 | |
| 183 | def decide_system_prompt(self, tar_lang="zh", multi_images_input=False): |
| 184 | zh = tar_lang == "zh" |
| 185 | self.is_vl |= multi_images_input |
| 186 | task_type = zh + (self.is_vl << 1) + (multi_images_input << 2) |
| 187 | return SYSTEM_PROMPT_TYPES[task_type] |
| 188 | |
| 189 | def __call__(self, |
| 190 | prompt, |
| 191 | system_prompt=None, |
| 192 | tar_lang="zh", |
| 193 | image=None, |
| 194 | seed=-1, |
| 195 | *args, |
| 196 | **kwargs): |
| 197 | if system_prompt is None: |
| 198 | system_prompt = self.decide_system_prompt( |
| 199 | tar_lang=tar_lang, |
| 200 | multi_images_input=isinstance(image, (list, tuple)) and |
| 201 | len(image) > 1) |
| 202 | if seed < 0: |
| 203 | seed = random.randint(0, sys.maxsize) |
| 204 | if image is not None and self.is_vl: |
| 205 | return self.extend_with_img( |
| 206 | prompt, system_prompt, image=image, seed=seed, *args, **kwargs) |
| 207 | elif not self.is_vl: |
| 208 | return self.extend(prompt, system_prompt, seed, *args, **kwargs) |
| 209 | else: |
| 210 | raise NotImplementedError |
| 211 | |
| 212 | |
| 213 | class DashScopePromptExpander(PromptExpander): |
nothing calls this directly
no outgoing calls
no test coverage detected