MCPcopy Create free account
hub / github.com/microsoft/BitNet / generate_all

Method generate_all

gpu/generate.py:217–304  ·  view source on GitHub ↗
(
        self, prompts: list[list[int]], use_cuda_graphs: bool, use_sampling: bool
    )

Source from the content-addressed store, hash-verified

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

Callers 1

mainFunction · 0.80

Calls 3

phaseMethod · 0.95
end_phaseMethod · 0.95
StatsClass · 0.90

Tested by

no test coverage detected