perf: project only sampled rows through lm_head during prefill

- add logits_positions to AutoRegressiveLM.forward, gathering rows before the final norm so the lm_head GEMM covers only the positions prefill samples from
- execute_prefill builds last_token_indices up front and passes them in, dropping the post-forward gather of a [tokens, vocab] tensor
- prefill graph warmup passes a single index; decode stays untouched (every row is sampled) and prefill itself runs eager, so graph capture is unaffected
- update the ragged-prefill fake to slice by the received index and add a packed-row exact-equality test

Benchmark: NVIDIA L20 (idle), CUDA 12.8, torch 2.11.0+cu128, 1.2B bf16 checkpoint, 512-token prompts, greedy; prefill B=32: 368.1 -> 323.5 ms (44.5k -> 50.6k tok/s, +13.8%), B=8: 89.3 -> 78.9 ms (+13.2%), B=1: 12.3 -> 11.4 ms (+7.9%); decode step unchanged; full suite: 897 passed
This commit is contained in:
2026-09-05 00:22:15 +08:00
parent 074642b6d2
commit a77e35dd51
4 changed files with 69 additions and 6 deletions
+10 -4
View File
@@ -131,6 +131,9 @@ def _warmup_cuda_graphs(
kv_cache=kv,
position_ids=pos_in,
fwd="prefill",
logits_positions=torch.tensor(
[warmup_len - 1], dtype=torch.long, device=dev
),
)
task_cache.task_free(tid)
@@ -346,6 +349,11 @@ class Executor:
]
)
# Last packed position per request; the model projects only these rows.
last_token_indices = (
torch.tensor(q_lens, dtype=torch.long, device=self.device).cumsum(0) - 1
)
with (
torch.inference_mode(),
timed(
@@ -363,11 +371,9 @@ class Executor:
start_pos=start_pos,
),
fwd="prefill",
logits_positions=last_token_indices,
)
last_token_indices = (
torch.tensor(q_lens, dtype=torch.long, device=self.device).cumsum(0) - 1
)
logits = outputs["logits"][last_token_indices]
logits = outputs["logits"]
step_out, _ = self._sample_logits(logits, tasks, return_logprobs)
return tasks, step_out