Modern AI Engineering

Lesson 4.12 · 25 min

Multi-Head Attention: Many Perspectives at Once

One attention pattern per token is a single opinion about what matters. What if each layer could hold a dozen different opinions at once, for the same cost?

In short: Multi-head attention runs several smaller attention operations ("heads") side by side. Each head has its own query, key and value projections, so it can learn its own pattern, such as "look at the previous word" or "look at the subject". The head outputs are concatenated and mixed by one more matrix, W_o. Because each head works in d_model / h dimensions, the total cost is about the same as one big head.

What is multi-head attention?

Multi-head attention (MHA) is the version of attention used inside every Transformer layer. Instead of computing one set of attention weights per token, it computes h separate sets in parallel, one per head. Each head is a complete, smaller copy of the attention mechanism with its own learned projection matrices. At the end, the heads' results are joined together and blended into a single output.

Think of it like a panel of specialist editors Hand a sentence to one editor and they read it with one focus. Hand it to eight editors, one checking grammar, one tracking who "she" refers to, one checking dates, one following the topic, and you get a richer review. Each editor writes short notes; the chief editor (W_o) combines them into one report. The heads are the specialists.

A quick recap of self-attention

In self-attention, every token of a sequence attends to the tokens of the same sequence. From the token matrix X (tokens × d_model) we compute queries Q = X·W_q, keys K = X·W_k and values V = X·W_v. Scores Q·Kᵀ are divided by √dₖ, softmax turns each row into weights, and each token's output is the weighted average of the value vectors.

The key limitation: for each token, a single softmax produces one set of weights. Softmax tends to concentrate weight on a few tokens, so one head usually can only focus strongly on one or two relationships at a time. If the word "it" needs to know both which noun it refers to and which verb governs it, a single averaged view blurs those two needs together.

Why do we need multi-head attention?

Language has many relationships at once. In "The animal did not cross the street because it was too tired", useful links include: "it" → "animal" (what the pronoun refers to), "tired" → "it" (what is tired), "cross" → "street" (verb and object), and every word → its neighbour (local word order). A single attention pattern must compromise between these. Multiple heads remove the compromise: each head can specialize.

Researchers who inspect trained models do find heads with recognizable roles: heads that mostly look at the previous token, heads that link pronouns to nouns, heads that attend to punctuation or to the first token, and "induction heads" that help copy patterns seen earlier in the context. Not every head is that tidy, and studies have shown that many heads can be removed at test time with little loss, so heads overlap and are partly redundant. Still, having several gives the model room to represent different relationships in parallel.

Pause and think: If each head is just attention with different weights, why would different heads learn different things rather than all the same pattern?

They start from different random initial weights, and the output matrix W_o rewards them for contributing useful, non-duplicated information to the loss. Nothing forces them to be different, which is why some heads do end up redundant.

Step-by-step working of multi-head attention

From input to output

  1. Project: Compute Q = X·W_q, K = X·W_k, V = X·W_v with full d_model × d_model matrices. This is the same as applying h separate d_model × d_head projections, just packed into one matrix multiply.
  2. Split into heads: Reshape each of Q, K, V from (T × d_model) to (h × T × d_head): the first d_head columns belong to head 1, the next to head 2, and so on.
  3. Attend per head: Each head independently computes softmax(QᵢKᵢᵀ / √d_head)·Vᵢ, giving (T × d_head). All heads run in parallel as one batched operation.
  4. Concatenate: Place the h head outputs side by side to get back a (T × d_model) matrix.
  5. Mix with W_o: Multiply by W_o (d_model × d_model) so information from different heads can combine. The result goes to the residual connection and then the feed-forward network.

Same cost as one big head

A natural worry: if we have 8 heads, is attention 8 times more expensive? No. Each head works in a smaller space, d_head = d_model / h. With d_model = 512 and 8 heads, each head uses 64 dimensions. The four projection matrices are still d_model × d_model each, so the parameter count is 4·d_model², exactly what a single 512-wide head with an output projection would use. The score computations are split across heads, so the total arithmetic is also about the same.

Published configurations of well-known models
Modeld_modelHeads hd_head
Original Transformer (base, 2017)512864
GPT-2 small7681264
Llama 2 7B4,09632128

One thing that does grow with the number of heads is the number of attention maps: h maps of size T × T per layer. That matters for memory when sequences are long, and it is one reason later designs such as grouped-query attention share keys and values between heads (covered in a later lesson).

Pause and think: A model has d_model = 1024 and 16 heads. What is d_head, and what number do the scores get divided by?

d_head = 1024 / 16 = 64, and scores are divided by √64 = 8. The scale uses the per-head width, not d_model.

A simple example walk-through

Let us run a tiny multi-head layer: 4 tokens, d_model = 8, 2 heads of width 4. The weights are random (untrained), so the patterns mean nothing linguistically. The point is to watch the shapes and to see that the two heads produce different weight maps from the same input.

multi_head.py

import numpy as np
np.set_printoptions(precision=2, suppress=True)
rng = np.random.default_rng(7)
T, d_model, h = 4, 8, 2            # 4 tokens, model width 8, 2 heads
d_head = d_model // h              # each head works in 4 dims
X = rng.standard_normal((T, d_model))
W_q, W_k, W_v, W_o = (rng.standard_normal((d_model, d_model)) * 0.5 for _ in range(4))
def softmax(z):
z = z - z.max(axis=-1, keepdims=True)
e = np.exp(z)
return e / e.sum(axis=-1, keepdims=True)
def split_heads(M):                # (T, d_model) -> (h, T, d_head)
return M.reshape(T, h, d_head).transpose(1, 0, 2)
Q, K, V = split_heads(X @ W_q), split_heads(X @ W_k), split_heads(X @ W_v)
print("Q per head shape:", Q.shape)
scores = Q @ K.transpose(0, 2, 1) / np.sqrt(d_head)   # (h, T, T)
weights = softmax(scores)
for i in range(h):
print(f"head {i} attention weights:\n", weights[i])
heads = weights @ V                                   # (h, T, d_head)
concat = heads.transpose(1, 0, 2).reshape(T, d_model) # glue heads side by side
out = concat @ W_o                                    # mix heads together
print("concat shape:", concat.shape, "-> output shape:", out.shape)
print("params in W_q,W_k,W_v,W_o:", 4 * d_model * d_model)

Output:

Q per head shape: (2, 4, 4)
head 0 attention weights:
[[0.05 0.04 0.37 0.53]
[0.11 0.14 0.39 0.36]
[0.01 0.01 0.51 0.47]
[0.12 0.05 0.66 0.17]]
head 1 attention weights:
[[0.17 0.18 0.29 0.36]
[0.11 0.09 0.68 0.12]
[0.09 0.18 0.04 0.69]
[0.22 0.2  0.37 0.21]]
concat shape: (4, 8) -> output shape: (4, 8)
params in W_q,W_k,W_v,W_o: 256

Look at the first row of each head: head 0 sends 53% of token 1's attention to token 4, while head 1 spreads it more evenly. Even with random weights, different projections give different views. Training turns those arbitrary differences into useful specializations.

Single-head vs multi-head

Where multi-head attention is used and why it helps

Multi-head attention appears in every Transformer family:

  • Decoder-only LLMs (GPT, Llama, Claude, Mistral): multi-head causal self-attention in every layer.
  • Encoders (BERT and embedding models): multi-head bidirectional self-attention.
  • Encoder-decoder models (T5, translation models, Whisper): multi-head self-attention plus multi-head cross-attention (next lesson).
  • Vision Transformers: image patches play the role of tokens, and heads learn spatial relationships.

Advantages in summary: several relationships captured in parallel; no extra parameters versus one head of the same width; fully parallel on GPUs as one batched matrix operation; and some interpretability, because individual heads can be inspected.

Common mistakes and limits Reshaping in the wrong order (splitting tokens instead of features) silently mixes up heads; always split the feature dimension. Scaling by √d_model instead of √d_head. Forgetting W_o, which leaves heads unable to combine. And reading too much into a single head: head roles are fuzzy, overlapping and differ from model to model, so "head 5 does coreference" is a description of tendencies, not a rule.

Where this goes next Because every head stores its own keys and values during generation, multi-head attention is memory-hungry at inference. Multi-query and grouped-query attention keep many query heads but share key/value heads to shrink that memory. You will meet them in the efficiency module.

Worked example, step by step

We said one head must compromise between relationships. Let us put numbers on that. Take the token “it” and three tokens it could attend to: “animal”, “street” and “tired”. Suppose “it” needs two things at once: its referent (“animal”) and its property (“tired”). All scores below are illustrative, already scaled.

One head, then two

  1. One head tries to do both: It gives scores [2, 0, 2] to animal, street and tired. Softmax: e² ≈ 7.39, e⁰ = 1, e² ≈ 7.39, total 15.78. Weights ≈ [0.47, 0.06, 0.47].
  2. Push harder: Raise both wanted scores to 6: [6, 0, 6]. Weights ≈ [0.50, 0.00, 0.50]. The unwanted token is gone, but each wanted token is stuck at one half.
  3. See the limit: One softmax row is a budget of 100%. Two targets can never both get more than 50%. Three targets would be capped at 33% each.
  4. Two heads: Head A scores [3, 0, 0] and head B scores [0, 0, 3]. e³ ≈ 20.09, so each head gives its target 20.09 / 22.09 ≈ 0.91 and the other two tokens about 0.05 each.
  5. Keep the results apart: Head A's output fills the first half of the vector and head B's fills the second half. W_o can then combine “who” and “what state” instead of receiving one blurred average.
Weights of “it” on each token (illustrative scores, softmax computed exactly).
Setupanimalstreettired
One head, scores [2, 0, 2]0.470.060.47
One head, scores [6, 0, 6]0.500.000.50
Head A, scores [3, 0, 0]0.910.050.05
Head B, scores [0, 0, 3]0.050.050.91

The gain is not free. With two heads, each one sees only half of the vector's width, so each has less room to describe a token. That is the trade behind the choice of h: more heads means more separate budgets, but a narrower view for each. It is also why the head width in published models tends to stay at 64 or 128 while the number of heads grows with the model.

Practice: try it yourself

The most common multi-head bug is a reshape that has the right shape and the wrong content. We will reproduce it on purpose. We fill a matrix with labelled numbers, split it into heads the right way and the wrong way, and write a small test that tells them apart.

practice_split_heads.py

import numpy as np
T, d_model, h = 3, 4, 2            # 3 tokens, width 4, 2 heads
d_head = d_model // h
# Label every number so we can trace it: entry = 10 * token + feature
M = np.array([[10 * t + f for f in range(d_model)] for t in range(T)])
print("M (rows = tokens, columns = features):\n", M)
right = M.reshape(T, h, d_head).transpose(1, 0, 2)   # split the FEATURE axis
wrong = M.reshape(h, T, d_head)                      # same shape, wrong content
print("shapes:", right.shape, wrong.shape)
print("head 0, correct split:\n", right[0])
print("head 0, wrong split:\n", wrong[0])
# Test: head 0 must hold features 0 and 1 of EVERY token
def head0_ok(H):
return np.array_equal(H[0], M[:, :d_head])
print("correct split passes:", head0_ok(right))
print("wrong split passes  :", head0_ok(wrong))
# Round trip: gluing the heads back together must return M exactly
merged = right.transpose(1, 0, 2).reshape(T, d_model)
print("round trip ok:", np.array_equal(merged, M))

Output:

M (rows = tokens, columns = features):
[[ 0  1  2  3]
[10 11 12 13]
[20 21 22 23]]
shapes: (2, 3, 2) (2, 3, 2)
head 0, correct split:
[[ 0  1]
[10 11]
[20 21]]
head 0, wrong split:
[[ 0  1]
[ 2  3]
[10 11]]
correct split passes: True
wrong split passes  : False
round trip ok: True

Now change it:

  • Set h = 4, so each head is 1 number wide. Predict what right[0] prints before running.
  • Change right to M.reshape(T, h, d_head) with no transpose. Predict its shape, and say what the first axis now counts.
  • Set T, d_model, h = 2, 6, 3. Predict the three rows of “head 0, wrong split”.

Pause and think: In the wrong split, head 0's second row is [2, 3]. Attention would treat that row as a token. What would this head actually be comparing?

Pieces of the same token. [0, 1] and [2, 3] are the two halves of token 0's vector, and [10, 11] is half of token 1. The head would compute attention between fragments as if they were three tokens, so its weights no longer mean “token i looks at token j”. Nothing crashes, which is why only a content test with traceable numbers catches it.

Pause and think: A single head with scores [6, 0, 6] gave “animal” and “tired” 50% each. Could a cleverly trained single head give both of them 90%?

No. The weights in one softmax row always sum to 1, so two targets together can hold at most 100%, and 90% + 90% is impossible. To put high weight on two different tokens for two different reasons, the model needs two separate softmax rows, which means two heads (or two layers).

Key takeaways

  • Multi-head attention runs h smaller attention operations in parallel, each with its own projections.
  • Each head works in d_head = d_model / h dimensions, so total parameters and compute stay about the same as one head.
  • Heads can specialize in different relationships; their outputs are concatenated and mixed by W_o.
  • Shapes: (T × d_model) → split to (h × T × d_head) → attend → concat back to (T × d_model).
  • Head roles are tendencies, not rules, and many heads are partly redundant.

Key terms

  • Head: One independent attention computation with its own query, key and value projections.
  • d_head: The width of each head, usually d_model divided by the number of heads.
  • Concatenation: Placing the head outputs side by side to rebuild a d_model-wide vector.
  • Output projection (W_o): The learned matrix that mixes the concatenated head outputs.
  • Self-attention: Attention where queries, keys and values all come from the same sequence.
  • Induction head: A kind of head found in trained models that helps continue patterns seen earlier in the context.

← 4.11 Causal Masking: Preventing the Model from Seeing the Future · 4.13 Cross-Attention: Connecting Encoder Output to the Decoder →