Multi-head attention
Multi-head attention is a transformer layer that runs h scaled dot-product attentions in parallel, each on its own learned projections of the queries, keys and values, then joins the h outputs and multiplies them by an output matrix WO.
Last updated: 07 Oct, 2026 · NumPy
One Scaled dot-product attention gives each token one weighted mix of the others. Multi-head attention gives it several, each free to look for a different relation.
Running several attention heads on the same words
One self-attention, with its own query, key and value weights, is one attention head. The idea of multi-head attention is to create several heads for the same words. The video uses the "Thinking Machines" illustration from Jay Alammar's The Illustrated Transformer: head 0 has its own weights W0Q, W0K and W0V, computes Q0, K0 and V0, then the scores, the softmax and the product with V, and gives Z0. Head 1 does the same with W1Q, W1K and W1V and gives Z1.
Each head can capture different dependencies. In "the cat sat on the mat", Z0 for cat might be dominated by sat and Z1 by mat. Together the heads expand the model's ability to focus on different positions of the tokens, and the model gets more information. No one assigns a relation to a head: it comes out of training, and many heads are hard to interpret.
Concatenating the heads and projecting with W^O
With h = 8 heads there are eight outputs, Z0 to Z7. The paper projects each head to dk = dv = dmodel/h = 512/8 = 64 numbers, so WiQ, WiK and WiV are 512 × 64 and every Zi is n × 64. Joining the eight side by side gives n × 512, and multiplying by WO (512 × 512), trained jointly with the rest of the model, gives the layer's output, again n × 512.
WO belongs to the attention sub-layer. It mixes the heads and returns a dmodel-sized vector, so the residual addition and the next layer receive vectors of the size they expect. After it come Add & Norm and then the separate Position-wise feed-forward network.
Why eight heads cost about the same as one
Each head works in 64 numbers instead of 512. Eight heads of 64 hold the same number of projection weights as one head of 512, and their score products add up to the same number of multiplications. In the paper's words, "due to the reduced dimension of each head, the total computational cost is similar to that of single-head attention with full dimensionality".
d_model, n = 512, 10 # a 10-token sentence
for h in (1, 8):
d_k = d_model // h
weights = h * 3 * d_model * d_k + h * d_k * d_model # W_Q, W_K, W_V per head + W_O
score_mults = h * n * n * d_k # multiplications in all the QK^T
print(f"h = {h}: d_k = {d_k:>3}, projection weights = {weights:,}, "
f"score multiplications = {score_mults:,}, attention matrices = {h} x ({n}, {n})")h = 1: d_k = 512, projection weights = 1,048,576, score multiplications = 51,200, attention matrices = 1 x (10, 10) h = 8: d_k = 64, projection weights = 1,048,576, score multiplications = 51,200, attention matrices = 8 x (10, 10)
- Both settings have 1,048,576 projection weights, 4 × 512 × 512.
- Both need 51,200 multiplications for the scores of a 10-token sentence: 8 × 10 × 10 × 64 = 1 × 10 × 10 × 512.
- What changes is the number of attention patterns: eight 10 × 10 weight matrices instead of one.
Splitting The cat sat across two heads
A small run on the board's embeddings, with dmodel = 4 and h = 2, so each head works in dk = 2 numbers. The weights are random with a fixed seed.
One head's projections
W_Q, W_K, W_V = (rng.normal(0, 1, (4, 2)) for _ in range(3)) # 4 -> 2 numbers
Z_i, w_i = attention(E @ W_Q, E @ W_K, E @ W_V) # (3, 2) outputJoining the heads
concat = np.concatenate(heads, axis=1) # two (3, 2) blocks -> (3, 4)
out = concat @ W_O # W_O is (4, 4): back to d_modelimport numpy as np
def attention(Q, K, V):
s = Q @ K.T / np.sqrt(K.shape[1])
w = np.exp(s) / np.exp(s).sum(axis=1, keepdims=True)
return w @ V, w
E = np.array([[1, 0, 1, 0], [0, 1, 0, 1], [1, 1, 1, 1]], float) # The, cat, sat
d_model, h = 4, 2
d_k = d_model // h # each head works in 2 numbers
rng = np.random.default_rng(42)
heads = []
for i in range(h):
W_Q, W_K, W_V = (rng.normal(0, 1, (d_model, d_k)) for _ in range(3)) # head i's own weights
Z_i, w_i = attention(E @ W_Q, E @ W_K, E @ W_V)
heads.append(Z_i)
print(f"head {i} attention weights:\n{np.round(w_i, 4)}")
concat = np.concatenate(heads, axis=1) # (3, 2) + (3, 2) -> (3, 4)
W_O = rng.normal(0, 0.5, (h * d_k, d_model)) # output projection
out = concat @ W_O
print("concat shape", concat.shape, "-> output shape", out.shape)
print("output:\n", np.round(out, 4))head 0 attention weights: [[0.6111 0.2431 0.1458] [0.1945 0.3723 0.4332] [0.4362 0.3321 0.2317]] head 1 attention weights: [[0.302 0.3639 0.3342] [0.2751 0.5487 0.1762] [0.2432 0.5844 0.1723]] concat shape (3, 4) -> output shape (3, 4) output: [[0.5163 0.6253 0.2151 1.2935] [0.8424 0.5407 0.3644 1.4098] [0.6986 0.5105 0.3054 1.2417]]
What the two heads attended to
- The two heads find different patterns in the same three tokens: in head 0 The attends most to itself (0.6111), in head 1 sat attends most to cat (0.5844).
- Each head returns (3, 2); joined they make (3, 4), and WO keeps that shape.
- Every row of every head's weights sums to 1, because each head is a full scaled dot-product attention.
Single-head vs multi-head attention
| One head (d_k = 512) | Eight heads (d_k = 64) | |
|---|---|---|
| Attention patterns per layer | 1 | 8 |
| Projection weights | 1,048,576 | 1,048,576 |
| Score multiplications | n² × 512 | 8 × n² × 64 = n² × 512 |
| What a token can mix in | one weighted average | eight, joined by Wᴼ |
Where you use multi-head attention
- Every attention sub-layer of the transformer: the paper uses h = 8 in encoder self-attention, masked decoder self-attention and cross-attention.
- Larger models use more heads: BERT-base has 12 heads of 64 numbers and BERT-large 16.
- Inspecting a trained model: tools such as bertviz draw each head's weights and show that different heads attend to different words.
Related
- Previous: Scaled dot-product attention
- Next: Position-wise feed-forward network
- Reference: Jay Alammar, The Illustrated Transformer
- Reference: Vaswani et al., Attention Is All You Need (2017)
- Change
h = 2toh = 4in the two-head example (dk becomes 1) and print each head's weights. - Add
16to the head counts in the cost example and check the weights and multiplications. - Replace
W_Owithnp.eye(4): the output then equals the concatenation.
Every expert started right here.