Context length and the quadratic cost
Be able to compute the memory and time cost of attention as a function of the sequence length.
Prerequisites
- DAttentionrequired
- DTime complexity and big-O notationrequired
Intuition
Attention compares every token with every other one. With T tokens that is T² pairs — double the context and the cost in the attention step quadruples.
| Context | Pairs | Relative |
|---|---|---|
| 2 k | 4 M | 1× |
| 8 k | 67 M | 16× |
| 32 k | 1.0 G | 256× |
| 128 k | 16 G | 4 096× |
But two costs have to be kept apart:
- Computation in the prefill: O(T²·d) — it grows quadratically.
- Memory during generation: the KV cache grows linearly in T, but still becomes the bottleneck in practice since it stays there for the whole session.
FlashAttention removed the quadratic memory cost for the attention matrix, but not the quadratic computation.
Formal
FLOPs per layer (approximately, the forward pass):
- The Q, K, V, O projections: — linear in T.
- The attention scores and : — quadratic.
- The MLP (4× expansion): — linear.
Attention only dominates once . For that means contexts below about 4 k tokens are dominated by the linear terms — which surprises many people. The quadratic term only bites in earnest at long contexts.
Memory:
- The KV cache: per sequence — linear in T, but multiplied by the batch.
- The attention matrix: per head if it is materialised. FlashAttention computes it in blocks in SRAM and never materialises it → memory.
Three routes to a long context: sparse or local attention (which changes what is computed), FlashAttention (which changes how it is computed), and RAG (which changes whether it needs computing at all). The last is usually the cheapest — fetch the 5 000 relevant tokens instead of feeding in 200 000.
Code
def cost(T, d=4096, layers=32, h_kv=8, d_head=128, bytes_per_number=2):
attn_flops = 2 * T**2 * d * layers
lin_flops = (4 * d**2 + 8 * d**2) * T * layers
kv_gb = 2 * layers * h_kv * d_head * T * bytes_per_number / 1e9
return {"T": T, "attn_TFLOPs": round(attn_flops / 1e12, 1),
"linear_TFLOPs": round(lin_flops / 1e12, 1),
"attn_share": round(attn_flops / (attn_flops + lin_flops), 2),
"kv_cache_GB": round(kv_gb, 2)}
for T in (2_048, 8_192, 32_768, 131_072):
print(cost(T))
# {'T': 2048, 'attn_TFLOPs': 1.1, 'linear_TFLOPs': 13.2, 'attn_share': 0.08, 'kv_cache_GB': 0.13}
# {'T': 8192, 'attn_TFLOPs': 17.6, 'linear_TFLOPs': 52.8, 'attn_share': 0.25, 'kv_cache_GB': 0.54}
# {'T': 32768, 'attn_TFLOPs': 281.5, 'linear_TFLOPs': 211.1,'attn_share': 0.57, 'kv_cache_GB': 2.15}
# {'T': 131072, 'attn_TFLOPs': 4504, 'linear_TFLOPs': 844, 'attn_share': 0.84, 'kv_cache_GB': 8.59}
The table shows the tipping point: at 2 k the attention is 8 % of the work, at 128 k it is 84 %.
Mastery means
- Computes the attention cost as a function of the sequence length
- Tells computation from memory
- Justifies why a long context is expensive
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Attention Is All You Need — arXiv (open access; licence per article)
- arXiv — FlashAttention — arXiv (open access; licence per article)