Attention Mechanics

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.

On this page 6
  1. Put numbers on it
  2. Why memory hurts before speed does
  3. The fix that changed everything
  4. What has not been fixed
  5. Remember this
  6. What to learn next

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.

TokensScores that must be computed
500250,000
1,0001,000,000
4,00016,000,000
32,0001,024,000,000
128,00016,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

Developer — Code and libraries.

Setup

bash
pip install numpy

What the growth actually looks like

quadratic.py
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}")
Output
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            64

Timings 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:

python
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

ApproachWhat it doesWhat it costs you
Sliding windoweach token sees only a nearby spandistant links must go through several layers
Sparse or block patternsa fixed subset of pairsthe pattern is chosen in advance, not by content
Low-rank or kernel methodsreplace softmax with a factorisable formquality gap widens at scale
Fewer key-value headsshrinks the cache, not the score matrixsmall 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

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

What to learn next