(
self,
prompt: str,
images,
max_gen_len: int,
temperature: float = 0.8,
top_p: float = 0.95,
modal = ['image'],
)
| 113 | |
| 114 | @torch.inference_mode() |
| 115 | def stream_generate( |
| 116 | self, |
| 117 | prompt: str, |
| 118 | images, |
| 119 | max_gen_len: int, |
| 120 | temperature: float = 0.8, |
| 121 | top_p: float = 0.95, |
| 122 | modal = ['image'], |
| 123 | ): |
| 124 | params = self.llma.params |
| 125 | |
| 126 | prompt_tokens = self.tokenizer.encode(prompt, bos=True, eos=False) |
| 127 | # truncate from the left. leave some space for generation. |
| 128 | max_seq_len = params.max_seq_len |
| 129 | if images is not None: |
| 130 | max_seq_len -= self.llma.image_words |
| 131 | |
| 132 | max_prompt_size = max_seq_len - max_gen_len |
| 133 | prompt_tokens = prompt_tokens[-max_prompt_size:] |
| 134 | |
| 135 | prompt_size = len(prompt_tokens) |
| 136 | |
| 137 | total_len = min(max_seq_len, max_gen_len + prompt_size) |
| 138 | |
| 139 | tokens = torch.full([total_len], 0).cuda().long() |
| 140 | |
| 141 | tokens[:len(prompt_tokens)] = torch.tensor(prompt_tokens).long() |
| 142 | start_pos = prompt_size |
| 143 | prev_pos = 0 |
| 144 | generate_until = start_pos |
| 145 | for cur_pos in range(start_pos, total_len): |
| 146 | logits = self.llma.forward_inference(tokens[None, prev_pos:cur_pos], prev_pos, images if prev_pos == 0 else None, modal = modal) |
| 147 | if temperature > 0: |
| 148 | probs = torch.softmax(logits / temperature, dim=-1) |
| 149 | next_token = self.sample_top_p(probs, top_p) |
| 150 | else: |
| 151 | next_token = torch.argmax(logits, dim=-1) |
| 152 | next_token = next_token.item() |
| 153 | |
| 154 | if next_token == self.tokenizer.eos_id: |
| 155 | break |
| 156 | |
| 157 | tokens[cur_pos] = next_token |
| 158 | prev_pos = cur_pos |
| 159 | generate_until = cur_pos + 1 |
| 160 | yield {"text": self.tokenizer.decode(tokens[start_pos:generate_until].tolist()), "end_of_content": False} |
| 161 | |
| 162 | yield {"text": self.tokenizer.decode(tokens[start_pos:generate_until].tolist()), "end_of_content": True} |
| 163 | |
| 164 | def sample_top_p(self, probs, p): |
| 165 | probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True) |
no test coverage detected