Attention is a differentiable dictionary lookup. A query is compared with every key, the comparison scores become weights through a softmax, and the output is the weighted average of the values. In a transformer this one operation does all the mixing of information between tokens, so its cost and numerical behaviour set the limits on context length, memory and decode speed.
This article treats attention as an algorithm. It covers the family of scoring functions, the exact reference algorithm and its complexity, how to make the softmax numerically safe, the online softmax that computes exact attention without ever storing the n by n score matrix, and incremental decoding with a cache. Each step comes with runnable NumPy code and a worked example you can check by hand. For intuition about why there are three projections and why scores are divided by the square root of d, see attention math and scaled dot-product attention.
A family of scoring functions
The idea is older than transformers. Nadaraya-Watson kernel regression (1964) already predicted a value as a similarity-weighted average of stored values. Neural attention made the similarity learnable. Three scoring functions have been widely used:
| Score s(q, k) | Introduced | Cost per pair | Notes |
|---|---|---|---|
| Additive: vT tanh(W q + U k) | Bahdanau et al., 2014 (translation) | O(da d) | Small MLP per pair; no single matmul, so slow at scale |
| Multiplicative: qT W k | Luong et al., 2015 | O(d2) or O(d) if W folded into k | Batches into a matrix product |
| Scaled dot product: qT k / sqrt(d) | Vaswani et al., 2017 | O(d) | One matmul for all pairs; scale keeps softmax out of saturation |
Dot-product scoring won because all n by m scores come from a single matrix product, which is the operation hardware does fastest. The rest of this article assumes it. Nothing below depends on the score being a dot product, though. Online softmax and incremental decoding work for any score you can compute one block at a time.
The reference algorithm and its cost
The reference algorithm for n queries, m keys and head dimension d has three steps. Compute S = Q KT / sqrt(d), apply a row-wise softmax to get P, and return O = P V. Masking is applied to S before the softmax by setting disallowed positions to minus infinity. It is never applied to P afterwards, because zeroing a weight after normalisation leaves rows that no longer sum to 1.
import numpy as np
def attention_ref(Q, K, V, mask=None):
"""Q: (n, d), K: (m, d), V: (m, dv), mask: (n, m) bool, True = allowed."""
S = Q @ K.T / np.sqrt(Q.shape[-1]) # (n, m) scores
if mask is not None:
S = np.where(mask, S, -np.inf)
m = S.max(axis=-1, keepdims=True) # row max, for stability
m = np.where(np.isfinite(m), m, 0.0) # fully masked row: avoid inf - inf
P = np.exp(S - m)
denom = P.sum(axis=-1, keepdims=True)
P = np.divide(P, denom, out=np.zeros_like(P), where=denom > 0)
return P @ V # (n, dv)Cost. Time is O(n m d) for the two matmuls plus O(n m) for the softmax. Memory is O(n m) for S and P. In self-attention n = m, so both are quadratic in sequence length. At n = 32,768 one fp32 score matrix for a single head is 4 GiB. That memory term, not the arithmetic, is what limited early long-context models, and the next two sections remove it.
Why subtract the max. exp(x) overflows fp32 above about 88 and fp16 above about 11. Softmax does not change if you subtract the same constant from every score in a row, so subtracting the row max puts every exponent at or below 0 and the largest term at exactly 1. The two guarded lines handle the edge case that breaks most hand-written implementations. When a row is fully masked, which happens with left padding, the max is minus infinity, inf minus inf is NaN, and that NaN spreads through every later layer. The reference code returns zeros for such rows, which is the usual convention.
Online softmax: exact attention without the score matrix
The softmax normaliser needs the whole row: you cannot divide by the sum until you have seen every score. The online softmax (Milakov and Gimelshein, 2018) gets around this. It keeps three running values per query: m, the largest score seen so far; l, the sum of exp(score minus m); and acc, the exp-weighted sum of values. When a new block of keys arrives and the max rises from m to m', every earlier term was computed with the wrong offset by a factor of exactly exp(m minus m'). So you multiply l and acc by that one correction factor and add the new block. At the end, acc divided by l is exactly the softmax-weighted sum. It is an exact reordering of the computation, not an approximation.
def online_attention(q, K, V, block=128):
"""Exact attention for one query q (d,), streaming over K/V in blocks."""
d = q.shape[-1]
m, l = -np.inf, 0.0
acc = np.zeros(V.shape[-1])
for start in range(0, K.shape[0], block):
s = K[start:start + block] @ q / np.sqrt(d) # scores for this block
m_new = max(m, s.max())
corr = np.exp(m - m_new) # exp(-inf) = 0 on the first block
p = np.exp(s - m_new)
l = corr * l + p.sum()
acc = corr * acc + p @ V[start:start + block]
m = m_new
return acc / lRun this over tiles of queries as well as tiles of keys and you have the outline of FlashAttention. The tiles live in on-chip SRAM, the n by m matrix is never written to memory, and memory falls to O(n d). The arithmetic is unchanged at O(n m d), so this is a memory and bandwidth optimisation, not a complexity one. For the GPU details see FlashAttention on GPUs. The same merge rule combines partial results computed on different devices, which is how ring and split-KV attention work: each partial carries (m, l, acc), and two partials merge exactly like two blocks.
Worked example: three keys, two blocks
Take d = 2, query q = (1, 0), and three keys k1 = (1, 0), k2 = (0, 1), k3 = (1, 1) with values v1 = (10, 0), v2 = (0, 10), v3 = (5, 5). The scaled scores are 1/sqrt(2) = 0.7071, 0 and 0.7071.
Reference. exp(0.7071) = 2.0281, so the weights are 2.0281, 1 and 2.0281 divided by their sum 5.0562, which gives 0.4011, 0.1978 and 0.4011. The output is 0.4011 x (10, 0) + 0.1978 x (0, 10) + 0.4011 x (5, 5) = (6.0167, 3.9833).
Online, blocks [k2] then [k1, k3]. After the first block, m = 0, l = 1 and acc = (0, 10). In the second block the max rises to 0.7071, so the correction factor is exp(0 - 0.7071) = 0.4931. The new terms are exp(0) = 1 for k1 and for k3. Then l = 0.4931 x 1 + 2 = 2.4931 and acc = 0.4931 x (0, 10) + (10, 0) + (5, 5) = (15, 9.9307). The output is acc / l = (6.0167, 3.9833), which matches the reference. Notice that the stale first block was corrected by a single multiplication. Nothing was recomputed.
Incremental decoding as memoisation
In autoregressive generation, token t attends to tokens 1 to t. Recomputing the keys and values of every previous token at every step costs O(t d2) in projections per layer per step, so generating n tokens costs O(n2 d2) in projections alone. The keys and values of past tokens never change, so you memoise them: project only the new token, append its k and v to a cache, and attend from one query.
class KVCache:
def __init__(self, d, dv):
self.K = np.empty((0, d))
self.V = np.empty((0, dv))
def step(self, x, Wq, Wk, Wv):
"""x: (d_model,) embedding of the newest token. Returns its attention output."""
q, k, v = x @ Wq, x @ Wk, x @ Wv
self.K = np.vstack([self.K, k]) # real engines preallocate or page this
self.V = np.vstack([self.V, v])
return online_attention(q, self.K, self.V)Each step now costs O(d2) for the projections plus O(t d) for attention, so attention for a full sequence is O(n2 d) in total. The cache trades memory, O(n d) per layer per head for K and V, for that saving. Causal masking also comes free, because the cache only holds the past. Per step, decode is a matrix-vector product over the whole cache, which means it is limited by memory bandwidth rather than compute. That is why engines work so hard on cache layout, sharing and compression.
Testing an implementation
Attention bugs rarely crash. They produce slightly wrong numbers that training then hides. Test against the reference using properties that must hold:
rng = np.random.default_rng(0)
Q, K, V = rng.normal(size=(4, 8)), rng.normal(size=(9, 8)), rng.normal(size=(9, 3))
ref = attention_ref(Q, K, V)
for b in (1, 2, 4, 9): # block size must not matter
got = np.stack([online_attention(q, K, V, block=b) for q in Q])
assert np.allclose(got, ref, atol=1e-10)
perm = rng.permutation(9) # keys are a set: order must not matter
assert np.allclose(attention_ref(Q, K[perm], V[perm]), ref)
big = attention_ref(Q * 1e3, K, V) # huge scores: no overflow
assert np.isfinite(big).all()
mask = np.zeros((4, 9), bool); mask[0, :] = True # rows 1-3 fully masked
assert np.allclose(attention_ref(Q, K, V, mask)[1:], 0.0)These checks run in milliseconds. They catch the common real defects: off-by-one block boundaries, a correction factor applied to l but not to acc, masking after the softmax, and NaNs from fully masked rows.
Failure modes
- Overflow in low precision. Computing exp in fp16 without subtracting the max gives inf. Accumulate m, l and acc in fp32 even when Q, K and V are bf16.
- NaN from empty rows. Padding plus a causal mask can leave a query with no allowed key. Guard the max and the division as the reference does.
- Masking with a large negative constant. Using -1e9 in fp16 overflows to -inf, while a value that is too small still leaks probability mass. Prefer real -inf with the guards above, or the dtype's minimum.
- Missing scale. Dropping the 1/sqrt(d) still trains, but the softmax saturates and gradients vanish for large d.
- Cache and position mismatch. Appending to the KV cache with the wrong position encoding gives fluent but wrong output. Test that cached decode equals full recomputation token by token.
Trade-offs: exact versus approximate
Exact attention is quadratic in time whatever you do. Online softmax removes the quadratic memory but not the quadratic arithmetic. Sub-quadratic alternatives change the algorithm. Sparse and sliding-window attention compute only some pairs. Linear attention replaces the softmax with a kernel feature map so the product can be reassociated. Low-rank methods compress keys. All of them give up exactness, and they tend to lose most on precise retrieval from far back in the context. Grouped-query attention, covered in grouped-query attention, takes a different path: it keeps attention exact and shrinks the KV cache by sharing keys and values across heads. The usual engineering order is exact attention with online softmax first, then cache reduction, and approximate attention only after measuring what it costs on your tasks.
What to do next
- Implement attention_ref and online_attention from this page and run the property tests, including the fully masked row.
- Work the three-key example by hand once, with the blocks in both orders, to see the correction factor act.
- Add a KV-cache decode loop and assert that it matches full recomputation for 50 random tokens.
- Measure the memory of the reference at n = 4k, 16k and 32k, then confirm that the online version is flat in n.
- Read FlashAttention's tiling with the (m, l, acc) state in mind. It is the same merge rule run on SRAM tiles.
- Before adopting an approximate attention variant, benchmark it on a long-range retrieval task, not only on perplexity.