Args: qas: A list of question answer pairs in the form of `[[q1, a1], [q2,a2], ... , [qn, None]]`. last answer should be None for generation. image: PIL Image for multi-modal understanding max_gen_len: generation hyper-param t
(self, qas: List[List[str]], image: Optional[Image.Image],
max_gen_len=512, temperature=0.1, top_p=0.5, seed=0)
| 9 | |
| 10 | class SPHINXModel(MetaModel): |
| 11 | def generate_response(self, qas: List[List[str]], image: Optional[Image.Image], |
| 12 | max_gen_len=512, temperature=0.1, top_p=0.5, seed=0) -> str: |
| 13 | """ |
| 14 | |
| 15 | Args: |
| 16 | qas: A list of question answer pairs in the form of `[[q1, a1], [q2,a2], ... , [qn, None]]`. |
| 17 | last answer should be None for generation. |
| 18 | image: PIL Image for multi-modal understanding |
| 19 | max_gen_len: generation hyper-param |
| 20 | temperature: generation hyper-param |
| 21 | top_p: generation hyper-param |
| 22 | seed: random seed |
| 23 | |
| 24 | Returns: |
| 25 | str: response |
| 26 | """ |
| 27 | # to avoid sampling inconsistency among model parallel workers |
| 28 | torch.manual_seed(seed) |
| 29 | np.random.seed(seed) |
| 30 | |
| 31 | if image is not None: |
| 32 | image = image.convert("RGB") |
| 33 | target_size = getattr(self.llma, 'image_size', 224) # 448 for SPHINX-1k, 224 for SPHINX |
| 34 | transform = get_transform("padded_resize", target_size) |
| 35 | image = transform(image).to(list(self.parameters())[0]) |
| 36 | |
| 37 | conv = default_conversation() |
| 38 | assert qas[-1][1] is None |
| 39 | |
| 40 | conv.load_qas(qas) |
| 41 | prompt = conv.get_prompt() |
| 42 | # print(prompt) |
| 43 | |
| 44 | # each turn of response ends with `conv_seq` |
| 45 | conv_sep = conv.response_end_signal |
| 46 | |
| 47 | # since MetaModel.generate is originally designed for batched inference, |
| 48 | # we need to form a batch of size 1 here |
| 49 | response = self.generate( |
| 50 | prompts=[prompt], |
| 51 | images=image.unsqueeze(0) if image is not None else None, |
| 52 | max_gen_len=max_gen_len, |
| 53 | temperature=temperature, |
| 54 | top_p=top_p, |
| 55 | additional_stop_symbols=[conv_sep] |
| 56 | )[0] |
| 57 | |
| 58 | return response |
nothing calls this directly
no test coverage detected