Multi-head attention in detail
Be able to implement multi-head attention with masking and a KV cache, and verify it against a reference library.
Prerequisites
- DAttentionrequired
- DTransformers — the architecturerequired
Intuition
Attention in one sentence: every token asks (Q) all the tokens (K) how relevant they are, and fetches a weighted sum of their content (V). The weights are softmax(QKᵀ/√d).
Multi-head: instead of one large attention, h small ones are run in parallel, each with its own Q/K/V projections of a lower dimension (d_model/h). One head can learn «the previous token», another «the subject of the sentence». The outputs are concatenated and projected.
The causal mask: during generation token t may only see tokens ≤ t. You set −∞ in the logits for the future before the softmax.
The KV cache: when token t+1 is generated, K and V for all the earlier tokens are unchanged — save them. Then every new token costs O(t) instead of O(t²). The memory for the cache (layers × heads × t × d) is often what limits the context length in practice.
Code
import numpy as np
def softmax(x, axis=-1):
x = x - x.max(axis=axis, keepdims=True); e = np.exp(x); return e / e.sum(axis=axis, keepdims=True)
def attention(Q, K, V, causal=False):
d = Q.shape[-1]
S = Q @ K.swapaxes(-1, -2) / np.sqrt(d) # (..., T, T)
if causal:
T = S.shape[-1]
S = np.where(np.tril(np.ones((T, T), bool)), S, -1e9)
return softmax(S) @ V
def multi_head(x, Wq, Wk, Wv, Wo, h, causal=True):
T, d = x.shape; dh = d // h
Q, K, V = x @ Wq, x @ Wk, x @ Wv # (T, d)
split = lambda M: M.reshape(T, h, dh).transpose(1, 0, 2) # (h, T, dh)
out = attention(split(Q), split(K), split(V), causal) # (h, T, dh)
return out.transpose(1, 0, 2).reshape(T, d) @ Wo
rng = np.random.default_rng(0); d, h, T = 8, 2, 5
W = [rng.normal(0, 0.3, (d, d)) for _ in range(4)]
x = rng.normal(size=(T, d))
y = multi_head(x, *W, h=h)
print(y.shape) # (5, 8)
# causality: change the last token → the first row is unchanged
x2 = x.copy(); x2[-1] += 1
print(np.allclose(multi_head(x2, *W, h=h)[0], y[0])) # True
Formal
. The scaling keeps the logits at a variance of about 1 when have a variance of 1 per component (the sum of products has variance ); without it the softmax saturates and the gradients die. Multi-head: , , with . The complexity: in time, in memory for the score matrix per head — FlashAttention avoids materialising it. With a KV cache, decoding step becomes a attention: .
Mastery means
- Implements scaled dot-product attention with a causal mask
- Splits into heads and merges them back correctly
- Explains the KV cache and why it makes generation linear per token
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Attention Is All You Need — arXiv (open access; licence per article)
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0
- arXiv — FlashAttention — arXiv (open access; licence per article)