Previous commit broadcast Q@K^T (saved 1GB/step).
This commit broadcasts scores@V too (saves 2GB/step).
Before: V expanded from [kv_h, padded_len, d] to [H, padded_len, d]
4×100K×256×4B → 24×100K×256×4B = 400MB → 2.4GB allocation
After: broadcast matmul at kv_h level
se: [kv_h, gqa, P, 1, part_sz] @ V: [kv_h, 1, P, part_sz, d]
→ [kv_h, gqa, P, 1, d] → reshape to [H, P, d]
V stays at kv_h size: 400MB (no 2.4GB allocation)
Total per-decode-step memory for 100K context:
Before all GQA opts: 3.6GB (K expansion + V expansion)
After: 600MB (6x total reduction from GQA ratio=6)
This is the CCCL insight applied: transform_reduce with a compound type.
Instead of expanding to full head count then reducing, keep the reduction
at the minimal group size and broadcast the grouping dimension.
Phase 1 rewrite:
Before: for p in range(195): torch.bmm(Q, K_partition_p)
After: scores = torch.bmm(Q, K_all) # ONE launch for all 100K tokens
scores_parts = scores.view(H, P, part_sz) # reshape, no copy
part_out = torch.bmm(scores_exp_flat, v_parts_flat) # ONE launch
195 Python→CUDA round-trips → 2 round-trips.
Architecture informed by CCCL:
- summary_statistics.cu: fuse (max, exp_sum, weighted_output) computation
into a single reduction pass over the data. We do this by computing
Q@K^T over the ENTIRE sequence in one bmm, then reshaping to partitions
for the softmax statistics — the data is only read once from HBM.
- block_reduce_warp_reductions.cuh: Phase 2 reduction combines partition
statistics using the same (rescale, accumulate) pattern as CUB's
cross-warp aggregate merging.
Phase 2 (unchanged, already vectorized):
global_max + rescale + torch.bmm(weights, partition_outputs)
Total GPU kernel launches per decode step:
Before: 1 (gather) + 195 (Q@K) + 195 (scores@V) + 1 (reduce) = 392
After: 1 (gather) + 1 (Q@K_all) + 1 (scores_exp@V) + 1 (reduce) = 4
KV gather also stays batched: key_cache[blk_ids] is one index_select.
The single biggest performance bottleneck in the baseline:
paged_attention_v2 = raise NotImplementedError()
paged_attn.py: use_v1 = True (hardcoded to avoid calling V2)
V1 limitation: processes entire KV sequence in one kernel launch.
For seq_len=100K, this is a single massive attention computation.
V2: splits into PARTITION_SIZE=512 chunks, runs them in parallel,
then reduces with log-sum-exp. 195 parallel partitions vs 1.
Implementation (paged_attention_v2_pytorch.py):
Phase 1: Per-partition attention
- For each (seq, head, partition): compute QK^T, softmax, weighted V sum
- Store partial: tmp_output, exp_sums, max_logits (per partition)
Phase 2: Cross-partition reduction (log-sum-exp)
- global_max = max(max_logits across partitions)
- rescale = exp(partition_max - global_max) × partition_exp_sum
- output = Σ (rescale / total_sum) × partition_output
This is the same algorithm as vllm's paged_attention_v2_kernel.cu:
- The reduction pattern is identical to CCCL's block_reduce_warp_reductions
(combine partial statistics from independent segments)
- The online softmax tiling is the same as Flash Attention's partitioning
Integration:
- patch_paged_attention_v2.py patches _custom_ops.py and paged_attn.py
- Removes use_v1=True hardcode → V2 used for seq_len > 8192
- Dockerfile adds the patch step
This is a PyTorch implementation (no CUDA compilation needed).
Next step: if /usr/local/corex/ has ixcc or nvcc-compatible compiler,
replace with compiled CUDA kernel for further speedup.