← All Posts
Deep Learning · Transformers· Attention and memory

The Quadratic Problem: Prefill, Decode, and a Complexity Explorer

“Attention is quadratic” is incomplete. Specify the operation, execution phase, and memory category. A full prompt pass, one cached decode step, an explicit score tensor, and a persistent KV cache have different scaling rules.

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.

Change the workload, inspect the assumptions
Estimates use the formulas in this chapter: batch one, two bytes per cached scalar, rectangular MHA attention arithmetic, and a GQA cache with adjustable KV head count. These are arithmetic and storage estimates, not measured latency.

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

ChangePrimary effectTerm that can remain
FlashAttentionAvoid full score materialization; improve IO scheduleDense pairwise arithmetic
GQA / MLAReduce cached representationGrowing token-indexed history
Sliding windowRestrict pairwise connectivityFinite-window retrieval tradeoff
Linear / delta recurrenceUse a state independent of prefix lengthState update cost and compressed-memory capacity limits
Hybrid stackUse fewer global-attention layersQuadratic prefill term from remaining global layers
Speculative decodingAmortize target verification across proposalsDraft 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.