Chapter 20
Multi-Head Attention
Heads as subspaces, reshapes and einsum, parameter and FLOP counts.
One attention head performs one soft lookup per token. Multi-head attention runs several such lookups in parallel, each in a lower-dimensional subspace, then mixes their results back into the model width. This is the attention layer used inside transformer blocks: projections create queries, keys, and values; heads attend; an output projection returns to the residual stream.
20.1 Projections and heads
Let the model width be and the number of heads be . Multi-head attention first applies four dense matrices:
The first three projections produce width arrays. They are reshaped into heads of width . Each head runs scaled dot-product attention from Chapter 19 with its own slice of , , and . The split is a reshape and transpose, not a learned operation; the learned part is the projection matrices around it. See Section B.5 for the matrix multiplication convention used throughout the code.
It is helpful to view each projection matrix as containing all heads side by side. The first columns of feed head one, the next columns feed head two, and so on. Training can therefore choose different query/key/value coordinates for different heads while keeping one dense matrix multiply per projection.
def split_heads(x, num_heads):
"""(B, T, d_model) -> (B, H, T, d_head)."""
batch, length, width = x.shape
if width % num_heads:
raise ValueError("model width must be divisible by the number of heads")
head_width = width // num_heads
return x.reshape(batch, length, num_heads, head_width).transpose(0, 2, 1, 3)
def combine_heads(x):
"""(B, H, T, d_head) -> (B, T, d_model)."""
batch, heads, length, head_width = x.shape
return x.transpose(0, 2, 1, 3).reshape(batch, length, heads * head_width)
The output projection is as important as the input projections. Without it, heads would remain separate blocks of features. With it, the model can mix evidence across heads before the result is added back to the residual stream.
This also means head identities are not sacred. A later layer sees only the mixed -wide output, not a labeled list of head decisions. Interpretability tools may name heads by behavior, but the computation is just linear projections, attention, concatenation, and another linear projection.
20.2 Forward pass
The NumPy forward pass is short because the previous chapter already owns single-head attention. After projecting, the code reshapes into , broadcasts the mask across heads, attends, combines heads, and applies .
The scale inside each head is , not , because each dot product uses only that head’s coordinates. Keeping the per-head score variance controlled lets the model add heads without changing the softmax temperature implied by width.
def multi_head_attention_forward(Xq, Xk, Xv, params, num_heads, mask=None,
return_cache=False):
"""Project, split into heads, attend, concatenate, and project out."""
W_Q, W_K, W_V, W_O = params
Q_linear, K_linear, V_linear = Xq @ W_Q, Xk @ W_K, Xv @ W_V
Q = split_heads(Q_linear, num_heads)
K = split_heads(K_linear, num_heads)
V = split_heads(V_linear, num_heads)
head_output, attn_cache = attention_forward(Q, K, V, _head_mask(mask), True)
joined = combine_heads(head_output)
output = joined @ W_O
if not return_cache:
return output
cache = (Xq, Xk, Xv, params, num_heads, attn_cache, head_output, joined)
return output, cache
Causal masking is shared across heads. If query position cannot attend to key position , no head may use that edge. The mask therefore broadcasts over the head axis and zeros the same future positions in every head’s attention matrix. Padding masks work the same way, except they usually depend on the example in the batch.
Several heads help because one attention distribution rarely serves every purpose. One head can focus on nearby syntax, another on a delimiter, and another on a long-range dependency. The projections let those heads form scores in different subspaces instead of forcing one set of query-key coordinates to explain all relations.
There is still no guarantee that heads specialize cleanly. Some heads may be redundant, especially in small models or late in training. The architectural bet is weaker and more useful: give the optimizer several independent attention distributions, then let the output projection decide how to combine them.
20.3 Backward pass
Backpropagation follows the forward graph in reverse. The output projection gives
The concatenation gradient is reshaped back into head blocks. Each head then uses the attention backward pass from Section 19.3, producing gradients for its query, key, and value slices. Combining those slices gives gradients for the projected arrays:
Input gradients are , and likewise for keys and values. In self-attention, the same feeds all three paths, so its total gradient is the sum of the query, key, and value input gradients. The tests finite-difference both that shared-input gradient and every projection matrix.
Cross-attention uses the same formulas but does not sum all three input paths into one array unless the caller actually reused the same input. Queries may belong to a decoder sequence while keys and values belong to an encoder sequence. The code therefore returns separate gradients for , , and .
def multi_head_attention_backward(grad_output, cache):
"""Backward pass for multi-head attention."""
Xq, Xk, Xv, params, num_heads, attn_cache, _head_output, joined = cache
W_Q, W_K, W_V, W_O = params
grad_W_O = _project_grad(joined, grad_output)
grad_joined = grad_output @ W_O.T
grad_heads = split_heads(grad_joined, num_heads)
grad_Q, grad_K, grad_V = attention_backward(grad_heads, attn_cache)
grad_Q_linear = combine_heads(grad_Q)
grad_K_linear = combine_heads(grad_K)
grad_V_linear = combine_heads(grad_V)
grad_Xq = grad_Q_linear @ W_Q.T
grad_Xk = grad_K_linear @ W_K.T
grad_Xv = grad_V_linear @ W_V.T
grad_W_Q = _project_grad(Xq, grad_Q_linear)
grad_W_K = _project_grad(Xk, grad_K_linear)
grad_W_V = _project_grad(Xv, grad_V_linear)
return (grad_Xq, grad_Xk, grad_Xv), (grad_W_Q, grad_W_K, grad_W_V, grad_W_O)
The transpose/reshape steps have no parameters, but they are still part of the derivative. Their backward pass is the inverse reshape/transpose, which is why the implementation reuses split_heads and combine_heads in reverse order.
20.4 Parameters, FLOPs, and variants
With matrices , , , and , standard multi-head attention has
parameters, independent of the number of heads as long as the total width stays fixed. For self-attention on a batch of sequences of length , the dominant multiply-add count scales as
The first term is the four dense projections. The second term is score computation and weighted value aggregation across all heads, because . More heads change the layout and the subspaces, not this leading-order total.
def parameter_count(model_width):
"""Four dense d_model by d_model matrices."""
return 4 * model_width * model_width
def self_attention_flops(batch, length, model_width):
"""Dominant multiply-add count: projections plus score/value products."""
projections = 4 * batch * length * model_width * model_width
attention = 2 * batch * length * length * model_width
return projections + attention
The term is why later chapters care about KV caches, grouped-query attention, and efficient kernels. Multi-head attention is expressive, but sequence length is expensive.
Memory has the same warning sign. The attention weights have shape in ordinary self-attention, so storing them for backward can dominate small educational implementations. Production kernels reduce the memory footprint by tiling or recomputing pieces, but the layer still represents interactions between pairs of positions.
|
In practice
|
The original Transformer used multi-head attention to let the model attend jointly to information from different representation subspaces [vaswani2017attention]. Modern decoder-only LLMs keep the same basic layer but often alter the key/value side for inference efficiency. Multi-query attention shares one set of keys and values across query heads [shazeer2019fast], and grouped-query attention shares them within groups of heads [ainslie2023gqa]. Those variants reduce KV-cache memory; the standard full multi-head version here is the clearest starting point. |
20.5 Teach it
The one-sentence version. Multi-head attention projects the same tokens into several query-key-value subspaces, runs attention in each, concatenates the results, and mixes them.
An analogy. Several readers skim the same paragraph with different highlighters; one marks names, one marks dates, one marks causes, and a final editor combines their notes.
At the board.
-
Draw entering , , and .
-
Split each projected width into blocks.
-
Run one attention equation per block with the same causal mask.
-
Concatenate blocks and multiply by .
Misconceptions to address.
-
"More heads always means more parameters." Not if is fixed.
-
"Heads see different tokens." They see the same allowed tokens through different projections.
-
"The mask is per head." Standard causal and padding masks broadcast to all heads.
Check for understanding. In self-attention, why must the input gradient add query, key, and value contributions?
20.6 Exercises
Explain why splitting attention into several heads can be more expressive than one head with the same total width.
Trace the shapes from through projection, split into heads, attention, concatenation, and output projection.
Derive the parameter count and the leading self-attention cost in (20.5).
Implement multi-head attention backward by calling the single-head attention backward per head, then finite-difference the shared self-attention input and all four projection matrices.
References
-
[ainslie2023gqa] J. Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. 2023. arXiv:2305.13245
-
[shazeer2019fast] N. Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019. arXiv:1911.02150
-
[vaswani2017attention] A. Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762