This project benchmarks three attention implementations at the decode step of autoregressive LLM generation, comparing standard SDPA against two variants of Flash-Decoding-style chunked attention across a range of KV cache lengths and batch sizes.
The central question is:
At what KV cache length does KV-parallel chunked attention outperform the standard fused SDPA kernel for q_len=1 decode steps?
FlashAttention is optimized for the prefill phase, where many query tokens can be processed in parallel. At decode time, only one new token is generated per step, so q_len=1. Standard FlashAttention parallelizes over query tokens — at q_len=1 there is nothing to parallelize over, leaving the GPU underutilized.
Flash Decoding addresses this by parallelizing over the KV cache dimension instead. The KV cache is split into chunks, each chunk is processed by a separate thread block, and the partial results are reduced to a final output. This allows the GPU to stay busy even at batch_size=1, q_len=1.
This benchmark implements both sequential and parallel variants of chunked decode in pure PyTorch and measures where each strategy beats the reference SDPA implementation.
Explicit q @ k^T -> softmax -> @ v in standard PyTorch. No kernel fusion. Baseline for correctness and overhead comparison.
torch.nn.functional.scaled_dot_product_attention. Uses PyTorch's fused attention kernel (FlashAttention-2 backend where available). This is the reference implementation and production default.
Flash-Decoding-style sequential chunked reduction. KV cache is split into chunks. Each chunk computes local softmax statistics (max, sum, weighted values) and an online reduction combines them. The loop runs sequentially in Python. Overhead from loop iteration dominates at short KV lengths.
Same algorithm but each chunk is launched on a separate CUDA stream, allowing all chunks to execute concurrently on the GPU. After all streams complete, a synchronization barrier merges the partial results. This models the parallelism of a real Flash Decoding CUDA kernel but still pays Python-level overhead for stream management.
Automatically selects chunk size based on KV length to reduce loop overhead while maintaining good parallelism. Two variants: sequential and parallel.
q_len=1 is the exact decode regime. At q_len=1, the attention computation is memory-bandwidth-bound rather than compute-bound. The GPU spends most time reading the KV cache, not doing matmul. This is the regime where KV-parallel chunking provides the most theoretical benefit.
A real Flash Decoding implementation requires a custom CUDA kernel with per-threadblock partial sums, warp-level reductions, and specialized memory access patterns. This project implements the algorithm in PyTorch to measure the algorithmic tradeoff without the kernel engineering overhead. The results show where the benefit comes from and why a custom kernel is needed to fully realize it.
The parallel_chunked implementation uses CUDA streams to approximate the concurrency of a real Flash Decoding kernel. Each chunk runs on a separate stream. This demonstrates the benefit of chunk parallelism while adding Python-level stream management overhead that a fused kernel would not have.
SDPA uses a fused kernel that reads the KV cache once and computes the full attention in one pass. For short KV lengths, the overhead of launching and synchronizing multiple chunks exceeds the benefit of parallelism. The crossover point appears at approximately 4-8K tokens on RTX 2070.
Hardware: NVIDIA GeForce RTX 2070 PyTorch: 2.13.0
auto_parallel vs SDPA speedup (bs=1, q_len=1): kv=512: 0.10x (stream overhead dominates) kv=1024: 0.23x kv=2048: 0.46x kv=4096: 0.95x (crossover) kv=8192: 1.38x (parallel chunked wins) kv=16384: 1.36x
auto_sequential vs SDPA speedup (bs=1, q_len=1): kv=512: 0.90x kv=16384: 1.13x (wins only at extreme KV lengths)
-
SDPA wins for short KV lengths. The fused kernel reads the KV cache in a single efficient pass. Chunking overhead dominates when the KV cache is small.
-
Parallel chunked decode wins above approximately 8K KV tokens. At kv=8192, parallel chunked achieves 1.38x speedup vs SDPA. This is the regime where parallelizing over the KV dimension pays off.
-
Sequential chunked decode only wins at extreme KV lengths. The Python loop overhead limits speedup to 1.13x even at kv=16384. This explains why Flash Decoding requires a fused CUDA kernel.
-
Batch size 4 does not benefit from chunked decode. With multiple requests batched, SDPA already achieves good GPU utilization. The KV-parallel benefit of Flash Decoding is most valuable at batch_size=1.
-
The crossover point at 4-8K KV tokens has practical implications. Modern LLMs with 128K context windows operate well into the Flash Decoding benefit regime. This explains why vLLM, SGLang, and FlashInfer implement fused Flash Decoding kernels for long-context serving.
-
The measured gains are a lower bound on real Flash Decoding. Python stream management and synchronization overhead prevent achieving the full parallel speedup. A fused CUDA kernel would move the crossover to shorter KV lengths and achieve higher peak speedup.
This benchmark does not model:
- A real fused Flash Decoding CUDA kernel
- Multi-query attention (MQA) or grouped-query attention (GQA)
- Causal masking effects on KV access patterns
- Memory pressure from large KV caches at model scale
- Batched-decode with heterogeneous sequence lengths
Results represent the algorithmic tradeoff in pure PyTorch. Production Flash Decoding gains (e.g., in FlashInfer or vLLM) are typically larger.
This project closes a pair on attention optimization:
flash-attention-benchmark: attention during prefill (compute-bound, q_len >> 1) flash-decoding-bench: attention during decode (bandwidth-bound, q_len = 1)
Together they cover the full attention optimization surface across both serving phases.
Joao Felipe De Souza 2026