Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Why does multi-head attention exist? What advantages do multiple heads provide?
What you need to know
Why one head is not enough
In "The chef who trained in Lucknow cooked biryani", the token "cooked" needs to know its subject ("chef") and its object ("biryani"). A single softmax row must split its weight across both, producing a blurred average. With two heads, one can put most weight on "chef" and another on "biryani", and both signals survive.
The shapes
With d_model = 8, h = 2, so d_head = 4, and T = 4 tokens:
1import torch23B, T, d_model, h = 1, 4, 8, 24d_head = d_model // h # 45x = torch.randn(B, T, d_model)6W_qkv = torch.nn.Linear(d_model, 3 * d_model, bias=False)7W_o = torch.nn.Linear(d_model, d_model, bias=False)89q, k, v = W_qkv(x).split(d_model, dim=-1) # each (1, 4, 8)10split = lambda t: t.view(B, T, h, d_head).transpose(1, 2)11q, k, v = split(q), split(k), split(v) # each (1, 2, 4, 4)1213scores = q @ k.transpose(-2, -1) / d_head**0.5 # (1, 2, 4, 4): one T x T map per head14out = scores.softmax(dim=-1) @ v # (1, 2, 4, 4)15out = out.transpose(1, 2).reshape(B, T, d_model) # concat heads -> (1, 4, 8)16y = W_o(out) # mix heads -> (1, 4, 8)17print(scores.shape, y.shape)18# torch.Size([1, 2, 4, 4]) torch.Size([1, 4, 8])The key move is view(B, T, h, d_head): the 8 numbers per token are cut into 2 groups of 4. Each group gets its own 4 × 4 attention map. Concatenation puts them back side by side, and W_o lets the heads' outputs combine.
The cost is the same
Parameters for Q, K, V and O are 4 · d_model² whether you use 1 head or 12. For GPT-2 small (d_model = 768, 12 heads of 64), that is about 2.36 million per layer. What changes is how the work is divided, not how much work there is.
2026 practice: fewer K/V heads
Modern LLMs usually keep many query heads but share key/value heads. Multi-query attention (MQA) uses one K/V head for all queries; grouped-query attention (GQA) uses a few, for example 32 query heads sharing 8 K/V heads in Llama 3 8B. This shrinks the KV cache by 4× there, with little quality loss.
A real-life example
A code-completion assistant sees:
result = compute_total(items, tax=To predict the next token well, the model needs several facts at once: that it is inside an unclosed (, that tax is a keyword argument, and what numeric variables were defined earlier in the file. Interpretability studies have found heads with specialised behaviour of this kind — heads that attend to the previous token, and "induction heads" that find an earlier occurrence of the current token and copy what followed it. With one head, those needs would compete for a single distribution.
Follow-up questions to expect
- "Do more heads always help?" — No. With
d_modelfixed, more heads means smaller heads, and very small heads compare vectors poorly; most LLMs use 64 or 128 dimensions per head. Studies have also shown many heads can be pruned after training with little loss. - "Why is
W_oneeded?" — Without it, each head's output stays in its own slice of the vector;W_olets the model combine what different heads found. - "How do MQA and GQA differ from MHA?" — They share K and V across groups of query heads to cut KV-cache memory during generation.