The Quadratic Problem: Prefill, Decode, and a Complexity Explorer
Define the accounting unit
Let $n$ be sequence length, $d$ the concatenated query-head width, $H$ the number of heads, $d_h=d/H$, and $b$ bytes per stored scalar. For the first calculation assume batch one, equal Q/K/V head widths, dense MHA, and ignore biases. One multiply-accumulate (MAC) is counted as two floating-point operations (FLOPs). Real kernels also perform softmax, masking, normalization, and data movement.
Where the square comes from
In one layer, $QK^\top$ costs $n^2d$ MACs and multiplying weights by $V$ costs another $n^2d$. Thus dense rectangular attention contractions cost $2n^2d$ MACs, or $4n^2d$ FLOPs. Four dense Q/K/V/output projections add $4nd^2$ MACs. A conventional two-matrix MLP with intermediate width $4d$ adds $8nd^2$ MACs.
The pairwise contractions equal those projections plus that MLP when $2n^2d=12nd^2$, or $n=6d$, under these simplified assumptions. At shorter contexts, the word “quadratic” alone does not tell you which term dominates. A gated MLP or different head dimensions changes this comparison.
A causal mask permits only $n(n+1)/2$ query-key pairs. A kernel that skips masked tiles can save work relative to the rectangular count; a naive dense multiplication followed by masking still computes the full rectangle. The explorer below deliberately reports the rectangular convention so its formulas are unambiguous.
Separate temporary scores from persistent state
Materializing all attention scores for one layer requires $Hn^2$ scalars before accounting for probabilities or backward storage. FlashAttention avoids storing that full matrix in high-bandwidth memory while computing the same dense attention operation.
A KV cache stores vectors per token rather than per pair. Across $L$ layers with $H_{KV}$ key/value heads, it occupies approximately $2LnH_{KV}d_hb$ bytes for one sequence. GQA reduces $H_{KV}$; MLA changes the representation being stored; paging changes allocation and sharing. These address different parts of the budget.
Prefill processes known prompt tokens
All prompt tokens are available, so their projections and attention can use large matrix operations. The pass produces prompt representations and caches. Its attention contractions grow quadratically with prompt length for global attention. Time to first token also includes scheduling, prompt processing, and the final output projection; it is not simply a FLOP count divided by peak GPU FLOPs.
One decode step has one new query per head
With a valid cache of $n$ tokens, processing one additional token creates new Q/K/V and reads approximately $n+1$ keys and values. The attention contractions cost about $2(n+1)d$ MACs per layer, linear in the current prefix length. The token-wise projections and MLP process one row.
Generating $T$ tokens after a prompt of length $P$ accumulates attention reads proportional to $\sum_{t=1}^T(P+t)=PT+T(T+1)/2$. One cached step is linear in prefix length; the entire growing continuation still has a quadratic term in $T$. KV caching removes repeated prefix recomputation, not every dependence on history length.
One concrete scale
For $n=8192$, $d=4096$, and $H=32$, the two dense attention contractions require about 549.8 billion MACs per layer, or 1.10 trillion FLOPs. An explicit 32-head score tensor at two bytes per scalar occupies 4 GiB. With 32 layers, 8 KV heads, and $d_h=128$, the persistent GQA cache occupies 1 GiB per sequence.
These numbers refer to different scopes: arithmetic and scores for one layer, cache across 32 layers. Comparing them without their scopes would confuse transient and persistent costs. The explorer labels each quantity explicitly.
Why FLOPs do not directly predict time
A small decode batch can spend much of its time moving weights and KV data rather than saturating matrix multiplication. Larger batches reuse weights across more tokens but also increase live state and can affect latency. Prompt prefill often exposes larger matrix operations, but shapes, kernels, precision, and parallelism still matter.
A simple lower-bound perspective compares arithmetic time $F/\text{compute throughput}$ with data movement time $D/\text{memory bandwidth}$ and takes at least their maximum. Real execution also pays launch, synchronization, communication, and scheduling costs. Neither hardware peak should be substituted for achieved throughput without measurement.
Classify the proposed improvement
| Change | Primary effect | Term that can remain |
|---|---|---|
| FlashAttention | Avoid full score materialization; improve IO schedule | Dense pairwise arithmetic |
| GQA / MLA | Reduce cached representation | Growing token-indexed history |
| Sliding window | Restrict pairwise connectivity | Finite-window retrieval tradeoff |
| Linear / delta recurrence | Use a state independent of prefix length | State update cost and compressed-memory capacity limits |
| Hybrid stack | Use fewer global-attention layers | Quadratic prefill term from remaining global layers |
| Speculative decoding | Amortize target verification across proposals | Draft cost, rejection, target model work |
For concrete implementations, see FlashAttention, PagedAttention, and Transformers are RNNs. The counts above are derived from the stated shapes, not reported benchmarks from those papers.
Try it: Does FlashAttention make an ordinary global-attention model’s KV cache constant in sequence length?
No. Avoiding the temporary score matrix is separate from storing persistent per-token keys and values. The cache still grows unless the architecture or cache policy changes.
Verify an equation before benchmarking it
Use the math lab for small correctness checks. Then use the comparison map to choose which cost a particular architecture change is meant to address.