Two terms, two exponents

Take one transformer block: model width d, sequence length s, batch 1, multi-head attention, an FFN of width 4d. Counting a multiply-add as 2 FLOPs, the forward pass decomposes into terms that are linear in s and terms that are quadratic:

Q,K,V projections   3 · 2sd·d   =  6 s d^2
output projection        2sd·d   =  2 s d^2
FFN (two 4d matmuls)    2·2s·4d·d  = 16 s d^2
                                   ----------
linear total                        24 s d^2

QK^T  (scores)          2 s·s·d    =  2 s^2 d
A·V   (weighted sum)    2 s·s·d    =  2 s^2 d
                                   ----------
quadratic total                      4 s^2 d

Note what is not quadratic: the Q/K/V and output projections are per-token matmuls, so they grow linearly like the FFN does. Only the two operations that touch every pair of positions — forming the score matrix and applying it — carry s². Swapping the FFN for SwiGLU changes nothing structural: three matrices at d_ff = 8d/3 gives 2 · 3 · s · d · (8d/3) = 16 s d² again.

Advertisement

Where the crossover sits

Divide the two totals and the whole question collapses to one ratio:

quadratic / linear  =  4 s^2 d / (24 s d^2)  =  s / (6d)

crossover:  s* = 6d

Below s = 6d the model is dominated by its per-token matmuls; above it, by pairwise attention. The numbers are unintuitive. For d = 768 the crossover is s* = 4608. For d = 2048 it is 12,288. For a frontier-width d = 8192 it is 49,152 — which is why large models can run 32k contexts while spending most of their FLOPs on the FFN, and why the quadratic scare story is really a small-model problem.

Two cautions on reading this. First, s* is where the quadratic term equals the linear term, not where ‘attention is half the cost’ — attention including its own projections is about two-thirds of the block at that point. Second, grouped-query attention does not move s*: sharing KV heads shrinks the cache, but every query head still forms its own s × s score matrix.

Advertisement

Prefill: the quadratic term, paid all at once

Prefill is the forward pass over the prompt. All s tokens go through together, so the matmuls are large GEMMs with high arithmetic intensity and the phase is compute-bound. Its cost is the block formula times layers:

prefill FLOPs  ≈  L · (24 s d^2 + 4 s^2 d)
               =  2 · N · s     +  4 L d · s^2

The first form is the honest one; the second shows that the familiar 2Ns rule of thumb (N = non-embedding parameters) is only the linear half. Quoting 2Ns alone understates prefill by a factor of 1 + s/(6d) — a rounding error at 1k tokens, 4.6× at 32k for a 1.5k-wide model.

The practical consequence is that prefill is where long context hurts first and hurts visibly: it is a single blocking wait before the first token appears, and it is the term that grows superlinearly. Doubling the prompt does not double time-to-first-token; it roughly triples it once you are past s*.

Decode: linear per step, quadratic in aggregate

Decode is the opposite regime. Each step processes exactly one token, so every matmul is a matrix-vector product: roughly one FLOP per byte loaded, which makes decode memory-bandwidth-bound. Per generated token you re-read the entire weight matrix set plus the entire KV cache.

decode FLOPs/token  ≈  2N + 4 L d · s      (linear in s)
bytes read/token    ≈  weight_bytes + kv_bytes(s)

So a single decode step is only linear in context length. The quadratic reappears in aggregate: generating T tokens after a prompt of s costs Σ_t 4Ld(s+t) ≈ 4Ld(sT + T²/2), quadratic in the generation length. A 4k-token answer is not four times a 1k-token answer.

This asymmetry is why the two phases want different optimisations. Prefill responds to better FLOP throughput — fused kernels, quantised GEMM, more cores. Decode responds only to moving fewer bytes: smaller weights, smaller cache, fewer KV heads. Tuning the wrong one is the most common wasted effort in CPU inference work.

The KV cache in bytes

The cache holds one key and one value vector per layer per token. The exact size, with no hand-waving:

kv_bytes = 2 · b · s · L · n_kv · d_head · bytes_per_elem
     2 = one K plus one V
     n_kv · d_head = the KV width (= d under plain MHA)

Everything on the right is fixed at model-build time except b and s, so cache size is strictly linear in sequence length — the quadratic lives in compute, not here. That linearity is deceptive, though, because the constant is large. Reference model for the rest of this article: L = 28, d = 1536, 12 heads of d_head = 128, MHA, FFN 4d, giving N ≈ 12Ld² = 0.79B non-embedding parameters. In fp16 that is 2 × 28 × 1536 × 2 = 172,032 bytes — 168 KiB per token. Grouping to 2 KV heads divides that by 6; int8 quantisation halves it again.

The same model at four sequence lengths

Holding that reference model fixed and sweeping s makes the two growth rates concrete. Crossover is s* = 6d = 9216.

sKV cache (fp16, MHA)quad / linearprefill FLOPs
51288 MB0.060.86 TFLOP
2,048352 MB0.223.97 TFLOP
8,1921.41 GB0.8924.5 TFLOP
32,7685.64 GB3.56237 TFLOP

Read the columns against each other. Memory grows 64× across the sweep, exactly with s. Prefill FLOPs grow 276×. And at 32k the fp16 cache alone is 5.64 GB — more than ten times the 0.45 GB the weights occupy at 4 bits. Past the crossover the cache, not the model, is the memory footprint.