| 215 | |
| 216 | @torch.inference_mode() |
| 217 | def generate_all( |
| 218 | self, prompts: list[list[int]], use_cuda_graphs: bool, use_sampling: bool |
| 219 | ) -> Tuple[Stats, list[list[int]]]: |
| 220 | bs = len(prompts) |
| 221 | prompt_lens = [len(p) for p in prompts] |
| 222 | padded_prompt_lens = [self.gen_args.prompt_length] * bs |
| 223 | max_prompt_length = max(prompt_lens) |
| 224 | gen_length = self.gen_args.gen_length |
| 225 | max_seq_length = max_prompt_length + gen_length |
| 226 | print(max_prompt_length, gen_length) |
| 227 | |
| 228 | bias = AttnBias.from_seqlens( |
| 229 | q_seqlen=padded_prompt_lens, |
| 230 | kv_seqlen=prompt_lens, |
| 231 | kv_padding=max_seq_length, |
| 232 | ) |
| 233 | bias.q_seqinfo.to("cuda") |
| 234 | bias.k_seqinfo.to("cuda") |
| 235 | |
| 236 | # Input tensors to the cuda graph |
| 237 | kv_seqlen = bias.k_seqinfo.seqlen |
| 238 | prompts = [prompt + [1] * (self.gen_args.prompt_length - len(prompt)) for prompt in prompts] |
| 239 | tokens = torch.IntTensor(sum(prompts, [])).cuda() |
| 240 | out_tokens = torch.zeros((max_seq_length, bs), dtype=torch.int) |
| 241 | |
| 242 | stats = Stats() |
| 243 | torch.cuda.synchronize() |
| 244 | stats.phase("prefill" if use_cuda_graphs else "total") |
| 245 | # stats.phase("total") |
| 246 | |
| 247 | output = self._prefill_compile_model(tokens, None) |
| 248 | |
| 249 | logits = output[kv_seqlen - 1, :] |
| 250 | logits = logits.view(bs, self.model_args.vocab_size) |
| 251 | |
| 252 | if use_sampling: |
| 253 | temp = 0.7 |
| 254 | top_p = 0.95 |
| 255 | probs = torch.softmax(logits / temp, dim=-1) |
| 256 | next_token = sample_utils.top_p(probs, top_p) |
| 257 | else: |
| 258 | next_token = torch.argmax(logits, dim=-1) |
| 259 | |
| 260 | next_token = next_token.reshape(bs) |
| 261 | out_tokens[0, :] = next_token |
| 262 | |
| 263 | torch.cuda.synchronize() |
| 264 | stats.phase("decode" if use_cuda_graphs else "total") |
| 265 | |
| 266 | eos_id = self.tokenizer.eot_id |
| 267 | for niter in range(1, gen_length): |
| 268 | kv_seqlen.add_(kv_seqlen < max_seq_length) |
| 269 | output = self._generate_compile_model(next_token, kv_seqlen) |
| 270 | |
| 271 | logits = output.view(bs, self.model_args.vocab_size) |
| 272 | |
| 273 | if use_sampling: |
| 274 | temp = 0.7 |