Architecture document: docs/paged_attention_kernel_architecture.md
Defines every module from CCCL algorithm patterns before code.
Three-level decomposition from CCCL:
Level 1 (warp_reduce_shfl): shfl.down butterfly for per-thread QK scores
Level 2 (block_reduce_warp_reductions): warp partials → SMEM → block aggregate
Level 3 (agent_scan decoupled lookback): cross-partition combine
Compound type (from summary_statistics.cu):
attention_partial = (max_score, exp_sum, weighted_v[256])
combine(a, b) = online softmax rescaling (same math as Flash Attention)
Key design change: Grid on num_kv_heads, not num_heads.
Before: grid = (1, 24, 200) = 4800 blocks, KV loaded 6x redundantly
After: grid = (1, 4, 200) = 800 blocks, KV loaded once per kv_head
Each block computes GQA_RATIO=6 query heads with shared KV loads.
Reduces KV cache bandwidth by 6x (the GQA ratio).
SMEM budget verified:
K tile [32, 256] fp16 = 16KB
V tile [32, 256] fp16 = 16KB
Total = 32KB ≤ 48KB ✓
Phase 1 kernel: _partition_attn_kernel
Processes query heads sequentially within the GQA group
to minimize register pressure (6 × 256 = 1536 registers
too many if all loaded simultaneously).
Phase 2 kernel: _reduce_partitions_kernel
Also gridded on kv_heads, reduces all partitions for
GQA_RATIO heads per block.
This replaces the previous Triton V2 which was gridded on num_heads
and had no GQA awareness at the kernel level.
Bug: if l_i > 0 branch in Triton is invalid (compiled as constexpr).
Also: p = exp(scores - m_i_new) computed after m_i_new update was
using the wrong reference max (should subtract m_ij first, then rescale).
Fix: Adapted exactly from prefix_prefill.py's proven-correct pattern:
p = exp(scores - m_ij) # probs relative to chunk max
l_ij = sum(p) # chunk sum
m_i_new = max(m_i, m_ij) # new running max
alpha = exp(m_i - m_i_new) # old accumulator rescale
beta = exp(m_ij - m_i_new) # new chunk rescale
l_i_new = alpha*l_i + beta*l_ij
acc = acc*(alpha*l_i/l_i_new) + (p*beta/l_i_new) @ V
This is the Flash Attention online softmax tiling algorithm.
Same math as CCCL's parallel_reduce with compound accumulators.