Why attention costs grow with the square of length
Every token is scored against every token, so doubling the input roughly quadruples the attention work and the memory it needs.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Attention scores every word against every other word, so twice the words means about four times the work.
Picture a small wedding reception where everyone greets everyone. Ten guests is forty-five greetings. Add ten more guests and it is not ninety — it is one hundred and ninety.
You doubled the guests and roughly quadrupled the greetings. Nobody did anything wrong. It is what happens when every person has to meet every other person.
Attention has the same shape. Every token greets every token.
Put numbers on it
A token is a chunk of text, roughly three quarters of an English word.
| Tokens | Scores that must be computed |
|---|---|
| 500 | 250,000 |
| 1,000 | 1,000,000 |
| 4,000 | 16,000,000 |
| 32,000 | 1,024,000,000 |
| 128,000 | 16,384,000,000 |
Look at the jump from 32,000 to 128,000. The text got four times longer. The score count went up sixteen times.
And this is for one head, in one layer. A real model has dozens of heads and dozens of layers, and every one of them does this.
Why memory hurts before speed does
Modern chips are extremely fast at multiplication. What they are not good at is holding enormous grids of numbers.
That grid of scores has to exist somewhere. At 32,000 tokens it is over a billion numbers, per head, per layer. Multiply by the heads and layers and no graphics card on earth has room.
This is why long documents used to fail with an out-of-memory error rather than being slow. The arithmetic was affordable. The storage was not.
The fix that changed everything
In 2022 a technique called FlashAttention removed the storage problem. It computes attention in small tiles that never leave the chip's fast memory.
The idea is closer to washing dishes than to mathematics. You do not lay every plate on the counter at once. You wash a few, stack them, and move on. The full grid never exists in one place.
The answer that comes out is identical, digit for digit. Only the memory is different. The work still grows with the square of the length. The storage grows only in step with the length.
That single change is why models went from a few thousand tokens of context to hundreds of thousands.
What has not been fixed
The multiplication count is still quadratic. Nobody has removed that while keeping quality.
There are methods that look at fewer pairs — nearby words only, or a sample of distant ones. They are faster and they lose something. The honest summary is that this is an open problem with many partial answers and no settled winner.
So when you hear that a model handles a million tokens of context, two questions are worth asking. What does it cost per request? And how much of that context does it actually use well?
Remember this
- Every token is scored against every token, so the work grows with the square of the length.
- Storage was the binding limit, and tiled attention removed it without changing the answer.
- The arithmetic is still quadratic, and cutting it always costs something.
What to learn next
- Softmax overflow inside attention — the other thing that breaks at scale.
- Context window — what the limit means from the outside.
- Latency and throughput — turning these counts into serving decisions.
Developer — Code and libraries.
Setup
pip install numpyWhat the growth actually looks like
import time
import numpy as np
def softmax(x):
x = x - x.max(axis=-1, keepdims=True)
e = np.exp(x)
return e / e.sum(axis=-1, keepdims=True)
d = 64
rng = np.random.default_rng(0)
print("the score matrix alone, for ONE head, in float16:")
print(f"{'tokens T':>10} {'T*T scores':>15} {'memory':>12}")
for T in (512, 2_048, 8_192, 32_768, 131_072):
cells = T * T
mb = cells * 2 / 1024 / 1024 # 2 bytes per float16 number
unit = f"{mb*1024:,.0f} KB" if mb < 1 else (f"{mb:,.0f} MB" if mb < 1024 else f"{mb/1024:,.1f} GB")
print(f"{T:>10,} {cells:>15,} {unit:>12}")
print("\nwhere the work goes, per layer, d_model=4096 (matmul FLOPs, forward only):")
d_model = 4096
print(f"{'T':>8} {'projections':>16} {'attention':>16} {'attention share':>17}")
for T in (128, 512, 2_048, 8_192, 32_768):
proj = 8 * T * d_model * d_model # Q,K,V,O: 4 matmuls, 2 FLOPs per multiply-add
attn = 4 * T * T * d_model # scores + weighted sum
print(f"{T:>8,} {proj:>16,} {attn:>16,} {attn/(proj+attn):>16.1%}")
print("\nmeasured wall-clock on THIS machine (numpy, CPU, float32).")
print("your absolute numbers will differ; the RATIOS are the point.")
print(f"{'T':>8} {'seconds':>10} {'x vs T=256':>12} {'T^2 predicts':>13}")
base = None
for T in (256, 512, 1024, 2048):
Q = rng.normal(size=(T, d)).astype(np.float32)
K = rng.normal(size=(T, d)).astype(np.float32)
V = rng.normal(size=(T, d)).astype(np.float32)
t0 = time.perf_counter()
for _ in range(5):
out = softmax(Q @ K.T / np.sqrt(d)) @ V
dt = (time.perf_counter() - t0) / 5
base = base or dt
print(f"{T:>8,} {dt:>10.4f} {dt/base:>12.1f} {(T/256)**2:>13.0f}")the score matrix alone, for ONE head, in float16:
tokens T T*T scores memory
512 262,144 512 KB
2,048 4,194,304 8 MB
8,192 67,108,864 128 MB
32,768 1,073,741,824 2.0 GB
131,072 17,179,869,184 32.0 GB
where the work goes, per layer, d_model=4096 (matmul FLOPs, forward only):
T projections attention attention share
128 17,179,869,184 268,435,456 1.5%
512 68,719,476,736 4,294,967,296 5.9%
2,048 274,877,906,944 68,719,476,736 20.0%
8,192 1,099,511,627,776 1,099,511,627,776 50.0%
32,768 4,398,046,511,104 17,592,186,044,416 80.0%
measured wall-clock on THIS machine (numpy, CPU, float32).
your absolute numbers will differ; the RATIOS are the point.
T seconds x vs T=256 T^2 predicts
256 0.0014 1.0 1
512 0.0028 1.9 4
1,024 0.0066 4.6 16
2,048 0.0324 22.7 64Timings vary between runs and between machines. The numbers above are one run on one laptop CPU. Do not expect to reproduce them; expect to reproduce the shape.
The three tables, read in order
Memory is the real cliff. 512 KB at 512 tokens, 32 GB at 131,072. One head. One layer. Multiply by 32 heads and 32 layers and the naive approach is not slow, it is impossible.
The 8,192 row is exactly 50 percent, and that is not a coincidence. The projections cost eight times length times width squared. Attention costs four times length squared times width. Set them equal and the length works out to twice the model width. With a width of 4,096, that is 8,192 tokens. Below that number, your model spends most of its time on the weight matrices. Above it, on attention.
The measured timings grow more slowly than the square at first, then faster. From 256 to 512 tokens the prediction is 4x and the measurement is 1.9x. Small runs are dominated by fixed overheads. They are also dominated by the parts of the work that are linear in length. From 1,024 to 2,048 the measurement is 4.9x against a predicted 4x. The score matrix has outgrown the processor's cache. Every access now costs a trip to main memory.
That last point generalises. Once a quadratic buffer stops fitting in fast memory, measured cost rises faster than the FLOP count predicts.
Why FlashAttention is not a different algorithm
The trick is that softmax can be computed in one pass over blocks. Keep a running maximum and a running sum, and rescale earlier partial results as new blocks arrive. This is the online-softmax method of Milakov and Gimelshein (2018).
Because of that, attention can be tiled. Load a block of queries and a block of keys. Compute their scores in on-chip memory, fold them into a running output, discard them. The full score matrix is never written out.
- FLOPs: still quadratic. Slightly more, in fact, because of recomputation in the backward pass.
- Memory traffic: linear in length.
- Result: numerically equivalent, not an approximation.
In PyTorch you get this by calling the fused function rather than writing the matmuls yourself:
import torch.nn.functional as F
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)Written against PyTorch 2.5.1. The dispatcher picks a FlashAttention, memory-efficient, or math backend based on dtype, device and mask. On CPU and in float64 you will fall back to the math backend, which does build the full matrix.
The four ways people actually cut the cost
| Approach | What it does | What it costs you |
|---|---|---|
| Sliding window | each token sees only a nearby span | distant links must go through several layers |
| Sparse or block patterns | a fixed subset of pairs | the pattern is chosen in advance, not by content |
| Low-rank or kernel methods | replace softmax with a factorisable form | quality gap widens at scale |
| Fewer key-value heads | shrinks the cache, not the score matrix | small quality loss, large memory win |
The last row is the one nearly every production model uses. It is the only one with no serious downside. See multi-head attention.
Common mistakes
Quoting quadratic cost while ignoring the linear term. Below roughly twice the model width, the weight matrices dominate. Optimising attention on a model that serves 1,000-token requests is effort spent in the wrong place.
Forgetting the KV cache during generation. At generation time the score matrix is one row wide, so attention is linear per step. The quadratic term reappears as the total over all steps, and the cache size is what actually fills memory. See context window.
Benchmarking without warming up. The first call allocates buffers and picks a kernel. Time the second call onwards, or your ratios are noise.
Assuming a large advertised context is usable end to end. Retrieval accuracy in the middle of a long context is measurably worse than at the edges. Long context is a capacity, not a guarantee.
Try it yourself
Change the memory table to count a full model instead of one head. Multiply by 32 layers and 32 heads. Find the largest T that fits in 80 GB. Then work out the same number under FlashAttention. There the cost per token per layer is roughly the KV cache, not the score grid. The gap between those two numbers is the reason long context became possible.
What to learn next
- Softmax overflow inside attention — the other thing that breaks at scale.
- Context window — what the limit means from the outside.
- Latency and throughput — turning these counts into serving decisions.
Researcher — Mathematics and papers.
Cost, stated precisely
Per layer, forward pass, sequence length $T$, model width $d$, with $d_{\text{ff}} = 4d$:
$$ C_{\text{proj}} = 8 T d^2 \quad (\text{Q, K, V, O}), \qquad C_{\text{ffn}} = 16 T d^2, \qquad C_{\text{attn}} = 4 T^2 d $$
Attention overtakes the projections when $4T^2 d > 8Td^2$, that is $T > 2d$. Including the feedforward layer, it overtakes the whole block at $T > 6d$. For $d = 4096$: 8,192 and 24,576 tokens respectively.
Memory for explicit scores is $\Theta(T^2)$ per head per layer. This is what fails first, and by a wide margin.
IO-awareness is the actual contribution
Dao et al. (2022), FlashAttention, arXiv:2205.14135, analyse attention in a two-level memory model. Fast SRAM of size $M$ sits above slow HBM. Standard attention moves $\Theta(T^2 + Td)$ words between them. Tiled attention moves
$$ \Theta!\left( \frac{T^2 d^2}{M} \right) $$
words, which for realistic $M$ and $d$ is many times fewer. The FLOP count is unchanged. The backward pass recomputes scores from the stored softmax statistics, trading extra arithmetic for far less traffic.
The lower bound in the companion analysis shows no exact attention algorithm can do asymptotically better in this model. FlashAttention-2 (arXiv:2307.08691) improves work partitioning across warps. FlashAttention-3 (arXiv:2407.08608) exploits asynchronous copy and FP8 on Hopper-class hardware.
Sub-quadratic families
Fixed sparsity. Sparse Transformer (Child et al., 2019, arXiv:1904.10509) uses strided and local patterns at $O(T\sqrt{T})$. Longformer (Beltagy et al., 2020) combines sliding windows with a few global tokens at $O(Tw)$. BigBird (Zaheer et al., 2020) adds random links. It proves the resulting pattern retains universal approximation and Turing completeness. The practical force of that is limited, since constants and depth requirements are not addressed.
Low rank. Linformer (Wang et al., 2020, arXiv:2006.04768) projects keys and values to a fixed length $k$, giving $O(Tk)$. The projection is length-specific, which makes variable-length inference awkward.
Kernel and linear attention. Katharopoulos et al. (2020) and Performer (Choromanski et al., 2021, arXiv:2009.14794) replace $\exp(q \cdot k)$ with $\phi(q)^\top \phi(k)$. That lets $(\phi(K)^\top V)$ be computed first and reused. Cost becomes $O(Td^2)$, with a constant-size recurrent state at decode time.
State-space models. Mamba (Gu and Dao, 2023, arXiv:2312.00752) achieves linear scaling with an input-dependent selective state space. It is not an attention approximation. Hybrid stacks interleave a minority of full-attention layers with state-space layers. They consistently outperform pure state-space stacks on recall-heavy tasks. That is the clearest evidence that exact pairwise comparison does something linear methods do not replicate.
Tay et al. (2020), Long Range Arena, arXiv:2011.04006, remains the cautionary benchmark. Many efficient variants that report strong perplexity fall behind on precise long-range retrieval.
The quality dimension
Efficiency results are frequently reported without the retrieval evaluation that would expose their cost. Two references worth pairing with any long-context claim:
- Liu et al. (2023), Lost in the Middle, arXiv:2307.03172, documents a U-shaped accuracy curve. The variable is the position of relevant information within a long context.
- Needle-in-a-haystack style probes measure exact retrieval at depth. They routinely separate models with identical advertised context lengths.
Papers
- Dao et al., FlashAttention, 2022 — arxiv.org/abs/2205.14135
- Milakov and Gimelshein, Online normalizer calculation for softmax, 2018 — arxiv.org/abs/1805.02867
- Child et al., Generating Long Sequences with Sparse Transformers, 2019 — arxiv.org/abs/1904.10509
- Wang et al., Linformer, 2020 — arxiv.org/abs/2006.04768
- Tay et al., Long Range Arena, 2020 — arxiv.org/abs/2011.04006
- Gu and Dao, Mamba, 2023 — arxiv.org/abs/2312.00752
What to learn next
- Softmax overflow inside attention — the other thing that breaks at scale.
- Context window — what the limit means from the outside.
- Latency and throughput — turning these counts into serving decisions.