Multi-head attention
Split the model's width into several narrow heads, let each one look for something different, then glue their findings back together.
- 12 min read
- 3 reading levels
- Updated
Read these first
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Multi-head attention runs several small searches side by side. Each hunts for a different kind of relationship, and the findings are combined.
Think about going to see a flat with three friends before you rent it. One of them checks the taps and the drainage. One checks how much light the rooms get. One asks the neighbours about water supply.
They all walked through the same flat. Each came back with different notes. Over chai afterwards you pool everything and make one decision.
That is multi-head attention. Same sentence, several different sets of eyes, one combined result.
Why one search is not enough
Read this: "the keys that Priya left on the table were not hers".
Several different relationships are live in that one sentence at the same time.
- "hers" points back to a person, and the only person is Priya.
- "were" has to agree with "keys", which is plural, not with "table".
- "left" needs to know what was left and where.
A single search has one query per word, so it produces one set of shares. Ask it to track the person, the plural agreement and the action at once. It has to average all three into one blurred answer.
Give the sentence three separate searches and each can specialise. Nobody has to compromise.
The trick that makes it free
Here is the part that surprises people. Adding heads costs nothing extra.
A model has a fixed width — a fixed number of slots per word. Multi-head attention does not add slots. It divides the existing ones.
Twelve heads in a model with 768 slots means each head works with 64 slots. Not 768 each. Sixty-four each.
So eight heads and one head use exactly the same number of learned numbers. You are not buying more capacity. You are choosing to spend the same capacity on several narrow specialists instead of one generalist.
one word, 768 slots wide
┌────────────────────────────────────────────────┐
│ │
└────────────────────────────────────────────────┘
│
split into 12 pieces of 64 slots
│
┌────┐┌────┐┌────┐┌────┐┌────┐┌────┐ ... ┌────┐
│head││head││head││head││head││head│ │head│
│ 1 ││ 2 ││ 3 ││ 4 ││ 5 ││ 6 │ │ 12 │
└────┘└────┘└────┘└────┘└────┘└────┘ └────┘
each does its own complete search
│
glue the 12 results back into 768 slots
│
one final mixing step
│
┌────────────────────────────────────────────────┐
│ the word, updated │
└────────────────────────────────────────────────┘What the trade actually is
Narrow heads are worse at any single job than one wide head would be. Sixty-four slots hold less than 768.
But most relationships in language are narrow. Matching a pronoun to a person does not need the whole width. Tracking singular against plural needs even less.
So the trade is usually worth it, and there is a limit. Push the head count high enough and each head becomes too thin to be useful for anything. Real models sit between eight and about a hundred and twenty-eight heads, depending on their width.
The honest part
You will read that head 4 of layer 7 "is the pronoun head". Treat this carefully.
Some heads in some trained models do have a clear, repeatable job. Many do not. Most are a blur of several partial jobs, and the same head can behave differently on different sentences.
Naming heads is a useful research tool and a bad mental model. Do not expect to open a model and find twelve tidy specialists waiting for you.
Remember this
- Several narrow searches run side by side, each free to look for something different.
- Heads split the model's width, so they cost nothing extra.
- Their results are glued back together and mixed by one final learned step.
What to learn next
- Causal masking — the one change that turns this into a language model.
- Reshape, view and contiguity — the tensor skill every head-splitting bug comes down to.
- Transformers — where multi-head attention sits in the full architecture.
Developer — Code and libraries.
Setup
pip install torchMulti-head attention, written out and checked against PyTorch
The safest way to learn this is to reimplement nn.MultiheadAttention and prove the two agree.
import torch
import torch.nn as nn
import torch.nn.functional as F
torch.manual_seed(0)
d_model, n_heads, T = 8, 4, 6
d_head = d_model // n_heads # heads split the width, they do not add to it
ref = nn.MultiheadAttention(d_model, n_heads, batch_first=True, bias=True)
x = torch.randn(1, T, d_model)
# --- the same computation, written out by hand -------------------------------
Wq, Wk, Wv = ref.in_proj_weight.chunk(3, dim=0) # torch packs all three here
bq, bk, bv = ref.in_proj_bias.chunk(3, dim=0)
def split_heads(t): # (1,T,d_model) -> (1,H,T,d_head)
return t.view(1, T, n_heads, d_head).transpose(1, 2)
q = split_heads(x @ Wq.T + bq)
k = split_heads(x @ Wk.T + bk)
v = split_heads(x @ Wv.T + bv)
print("x ", tuple(x.shape))
print("q per head", tuple(q.shape), " <- (batch, heads, tokens, width per head)")
scores = q @ k.transpose(-2, -1) / d_head ** 0.5 # every head gets its own T x T map
A = scores.softmax(dim=-1)
ctx = A @ v # (1, H, T, d_head)
merged = ctx.transpose(1, 2).reshape(1, T, d_model) # glue the heads back together
mine = merged @ ref.out_proj.weight.T + ref.out_proj.bias
theirs, _ = ref(x, x, x, need_weights=False)
print("\nmax difference from nn.MultiheadAttention:", float((mine - theirs).abs().max()))
print("\nwhere token 0 looks, head by head:")
for h in range(n_heads):
row = A[0, h, 0]
print(f" head {h}: " + " ".join(f"{p:.3f}" for p in row))
print("\nhow different are the heads? pairwise mean |difference| of their maps:")
for a in range(n_heads):
for b in range(a + 1, n_heads):
print(f" head {a} vs head {b}: {float((A[0,a]-A[0,b]).abs().mean()):.3f}")
print("\nparameter count is identical either way:")
for h in (1, 2, 4, 8):
m = nn.MultiheadAttention(d_model, h, batch_first=True)
print(f" {h} head(s): {sum(p.numel() for p in m.parameters())} parameters")x (1, 6, 8) q per head (1, 4, 6, 2) <- (batch, heads, tokens, width per head) max difference from nn.MultiheadAttention: 2.9802322387695312e-08 where token 0 looks, head by head: head 0: 0.394 0.156 0.024 0.199 0.068 0.160 head 1: 0.028 0.018 0.475 0.162 0.205 0.111 head 2: 0.222 0.167 0.063 0.060 0.361 0.126 head 3: 0.347 0.166 0.152 0.109 0.136 0.090 how different are the heads? pairwise mean |difference| of their maps: head 0 vs head 1: 0.125 head 0 vs head 2: 0.077 head 0 vs head 3: 0.077 head 1 vs head 2: 0.096 head 1 vs head 3: 0.092 head 2 vs head 3: 0.060 parameter count is identical either way: 1 head(s): 288 parameters 2 head(s): 288 parameters 4 head(s): 288 parameters 8 head(s): 288 parameters
The three lines worth staring at
2.98e-08 is a float32 match. The hand-written version and PyTorch's own module agree to the precision of the number format. If you are reimplementing attention, this is the test that tells you when to stop debugging.
Head 0 puts 0.394 on token 0 while head 1 puts 0.028 there. Same input, same layer, same step. These heads are looking at genuinely different places, and no training has happened yet. The divergence comes purely from independent random initialisation.
288 parameters, whatever the head count. One head or eight, the module holds the same weights. This is the single most misunderstood fact about multi-head attention. A two-line experiment settles it.
The reshape is the whole implementation
Every multi-head implementation is one reshape and one transpose:
# (batch, tokens, d_model) -> (batch, heads, tokens, d_head)
t.view(B, T, H, d_head).transpose(1, 2)
# ...and back again
ctx.transpose(1, 2).reshape(B, T, d_model)The .transpose(1, 2) puts heads next to batch. The matmuls that follow then treat each head as an independent problem. The .reshape at the end is only valid because transpose was applied first. See reshape, view and contiguity for why the ordering matters.
In production, use the fused path
import torch.nn.functional as F
ctx = F.scaled_dot_product_attention(q, k, v, is_causal=True) # q,k,v: (B, H, T, d_head)Written against PyTorch 2.5.1; the signature is unchanged in the current 2.13 documentation. This dispatches to FlashAttention or a memory-efficient kernel and never builds the full score matrix. On long sequences that is the difference between running and running out of memory.
Sharing keys and values across heads
At generation time, keys and values are cached for every token produced so far. That cache, not the weights, is what fills GPU memory. Two variants shrink it:
| Design | Query heads | Key/value heads | Cache size |
|---|---|---|---|
| Multi-head | H | H | full |
| Grouped-query | H | G, where G divides H | H/G smaller |
| Multi-query | H | 1 | H times smaller |
Llama 3 8B uses 32 query heads and 8 key-value heads, a fourfold cache reduction. PyTorch exposes this through enable_gqa=True on scaled_dot_product_attention, documented as experimental and CUDA-only. The portable route remains an explicit repeat_interleave on the key and value heads.
Common mistakes
.view(B, H, T, d_head) instead of .view(B, T, H, d_head).transpose(1, 2). The first one interleaves tokens into the wrong heads. It runs without error and trains to nothing. Reshape with tokens still adjacent, then transpose.
A head count that does not divide the model width. nn.MultiheadAttention raises immediately, which is a kindness. Hand-written code often silently truncates.
Calling .reshape after .transpose and expecting no copy. The tensor is no longer contiguous, so .reshape copies. .view would raise. Both behaviours are correct; know which one you are getting.
Reading need_weights=True output as head-by-head maps. By default nn.MultiheadAttention averages across heads before returning. Pass average_attn_weights=False to see individual heads.
Try it yourself
Change n_heads to 8 with d_model still 8, so each head has a single slot. Re-run and look at the attention rows. With one slot per head, every score is a product of two numbers. The heads become far more similar. That is the head-count limit, visible in six lines.
What to learn next
- Causal masking — the one change that turns this into a language model.
- Reshape, view and contiguity — the tensor skill every head-splitting bug comes down to.
- Transformers — where multi-head attention sits in the full architecture.
Researcher — Mathematics and papers.
Definition
$$ \operatorname{MultiHead}(X) = \operatorname{Concat}(\text{head}_1, \dots, \text{head}_h) W_O $$
$$ \text{head}_i = \operatorname{Attention}(X W_Q^{(i)}, X W_K^{(i)}, X W_V^{(i)}) $$
with $W_Q^{(i)}, W_K^{(i)} \in \mathbb{R}^{d_{\text{model}} \times d_k}$, $W_V^{(i)} \in \mathbb{R}^{d_{\text{model}} \times d_v}$ and $W_O \in \mathbb{R}^{h d_v \times d_{\text{model}}}$. Vaswani et al. set $d_k = d_v = d_{\text{model}} / h$, which makes the total parameter count $4 d_{\text{model}}^2$ independent of $h$.
Concatenation plus a shared projection is a sum of per-head projections
Partition $W_O$ row-wise into blocks $W_O^{(i)} \in \mathbb{R}^{d_v \times d_{\text{model}}}$. Then
$$ \operatorname{MultiHead}(X) = \sum_{i=1}^{h} \text{head}_i \, W_O^{(i)} $$
The heads do not interact inside the layer at all. Each writes independently into the residual stream, and the sum is the only place their outputs meet. This is what licenses per-head ablation and per-head circuit analysis in Elhage et al. (2021).
It also means the OV circuit of head $i$ is the product $W_V^{(i)} W_O^{(i)}$. That is a rank-$d_v$ map from the stream to itself.
Head width, not head count, is the binding constraint
Bhojanapalli et al. (2020), arXiv:2002.07028, study this directly. Fixing $d_k = d_{\text{model}} / h$ makes the per-head score matrix rank-limited at $d_k$. That rank, not $h$, determines which attention patterns are representable. Decoupling head width from $d_{\text{model}} / h$ keeps heads wide as $h$ grows. That improves quality at fixed $d_{\text{model}}$, at the cost of more parameters.
How much redundancy is there
Michel et al. (2019), arXiv:1905.10650, prune trained heads greedily, scoring importance by gradient. Many layers tolerate removal of most heads with small quality loss. A minority are load-bearing, and removing those is catastrophic. Voita et al. (2019), arXiv:1905.09418, reach the same conclusion via $L_0$ regularisation. They characterise the surviving heads as positional, syntactic or rare-token specialists.
The methodological caution matters as much as the finding. Importance is measured after training. It does not follow that the pruned heads were unnecessary during training.
Key-value sharing
Per-token KV cache size for one layer is $2 \cdot g \cdot d_k \cdot \text{bytes}$, where $g$ is the number of key-value groups. At long context this dominates inference memory.
- Multi-query attention (Shazeer, 2019, arXiv:1911.02150): $g = 1$. Large cache reduction; measurable quality loss and reported training instability at scale.
- Grouped-query attention (Ainslie et al., 2023, arXiv:2305.13245): $1 < g < h$. The paper also shows an existing multi-head checkpoint can be uptrained into GQA. Mean-pool key and value heads within each group, at roughly five percent of original pretraining compute.
- Multi-head latent attention, from the DeepSeek-V2 report (arXiv:2405.04434), caches a low-rank latent per token. It projects up to $K$ and $V$ at use time, trading arithmetic for cache bandwidth.
Sparsity across heads
Not every head needs the full sequence. Mixed designs are now common. A fraction of layers run sliding-window attention, restricting each query to a fixed nearby span. The remainder stay global. Longformer (Beltagy et al., 2020, arXiv:2004.05150) and BigBird (Zaheer et al., 2020, arXiv:2007.14062) established the pattern. Mistral 7B (arXiv:2310.06825) applied it throughout a decoder-only language model. The design space here is genuinely unsettled. Open models released through 2025 and 2026 disagree with each other. They differ on window size, on which layers get windows, and on how windows combine with grouped-query attention.
Papers
- Vaswani et al., Attention Is All You Need, 2017 — arxiv.org/abs/1706.03762
- Michel et al., Are Sixteen Heads Really Better Than One?, 2019 — arxiv.org/abs/1905.10650
- Voita et al., Analyzing Multi-Head Self-Attention, 2019 — arxiv.org/abs/1905.09418
- Bhojanapalli et al., Low-Rank Bottleneck in Multi-head Attention Models, 2020 — arxiv.org/abs/2002.07028
- Ainslie et al., GQA, 2023 — arxiv.org/abs/2305.13245
What to learn next
- Causal masking — the one change that turns this into a language model.
- Reshape, view and contiguity — the tensor skill every head-splitting bug comes down to.
- Transformers — where multi-head attention sits in the full architecture.