onnxruntime
e090fbf3 - LightningIndexer: score only the cache rows the step can reach

Commit
59 days ago
LightningIndexer: score only the cache rows the step can reach The scoring GEMM was sized by `past_cache->Shape()[1]`, which is the export's static max_seq_len / compress_ratio, and never consulted past_lens. The select kernel then discarded every row above each query's own `visible` bound. So a 256K-capable export serving a 512-token prompt scored 128 live rows against a 65,664-row capacity and threw away 99.8% of the result. That was not a rounding error in the budget. Profiling one prefill Run on DeepSeek-V4-Flash put 21 launches of this GEMM at 7.15 ms each -- 150 of 334 ms, the largest single kernel in prefill, with the indexer's own kernels taking the total to 59%. It also explains why prefill throughput was nearly flat in context: the dominant term already paid the maximum at every length. `score_capacity` now bounds the GEMM, the reduce grid and both scratch buffers. The bound is max_b(visible(past_lens[b], seq_len - 1)) + 1, which is the largest row any query in the step can select, so nothing that could have been read is skipped and the result is unchanged by construction. Two constraints shaped it. cuBLAS takes `m` on the host while past_lens is on the device, so the bound costs one synchronised copy per node -- negligible against the GEMM it removes. And it is taken only outside stream capture: a captured launch would freeze the bound at capture time and then replay it against a longer sequence, so graph-captured decode keeps the full capacity and is bit-for-bit the old path. The bound is rounded up to a power of two before it reaches the allocator. An exact bound grows a few rows per prefill chunk, so BFC saw a slightly different scratch size every chunk and stranded a block each time; that fragmented the arena badly enough to fail a 1.8 GB request on a 64K prompt where the unclamped 8.6 GB request had succeeded. Measured on 8xH200, DeepSeek-V4-Flash, max_seq_len=262656, prefill chunk 512: prefill goes 1472 -> 3261 tok/s at 1K (2.21x), 1385 -> 2853 at 4K, 1343 -> 2647 at 16K, 1290 -> 2310 at 64K and 1106 -> 1484 at 256K, where early chunks still have a short past. 256K TTFT falls 237.0 s -> 176.7 s. Correctness: greedy tokens are bit-identical to the previous build at 1K, 16K and 64K including tokens_per_step under speculative decode, MMLU-Pro 800 is 532/800 against 532/800 with zero discordant items, and the operator's parity test gains a case with capacity 65,538 against a few hundred live rows -- without it the clamp is never exercised and a bad bound would read as a pass.
Author
Committer
Parents
Loading