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 dd and the number of heads be HH. Multi-head attention first applies four dense matrices:

Q=XqWQ,K=XkWK,V=XvWV,Y=concat⁡(O1,…,OH)WO.(20.1)\mQ=\mX_q\mW_Q,\quad \mK=\mX_k\mW_K,\quad \mV=\mX_v\mW_V,\quad \mY=\operatorname{concat}(\mO_1,\ldots,\mO_H)\mW_O .\tag{20.1}

The first three projections produce width dd arrays. They are reshaped into HH heads of width dh=d/Hd_h=d/H. Each head runs scaled dot-product attention from Chapter 19 with its own slice of Q\mQ, K\mK, and V\mV. 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 dhd_h columns of WQ\mW_Q feed head one, the next dhd_h 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.

Listing 20.1 Splitting and combining heads
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 WO\mW_O 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 dd-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 (B,T,d)(B,T,d) into (B,H,T,dh)(B,H,T,d_h), broadcasts the mask across heads, attends, combines heads, and applies WO\mW_O.

The scale inside each head is dh\sqrt{d_h}, not d\sqrt{d}, 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.

Listing 20.2 Multi-head attention forward pass
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 tt cannot attend to key position uu, 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

WˉO=Ycat⊤Yˉ,Yˉcat=YˉWO⊤.(20.2)\bar{\mW}_O=\mY_{\text{cat}}^\T\bar{\mY},\qquad \bar{\mY}_{\text{cat}}=\bar{\mY}\mW_O^\T .\tag{20.2}

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:

WˉQ=Xq⊤Qˉ,WˉK=Xk⊤Kˉ,WˉV=Xv⊤Vˉ.(20.3)\bar{\mW}_Q=\mX_q^\T\bar{\mQ},\quad \bar{\mW}_K=\mX_k^\T\bar{\mK},\quad \bar{\mW}_V=\mX_v^\T\bar{\mV}.\tag{20.3}

Input gradients are Xˉq=QˉWQ⊤\bar{\mX}_q=\bar{\mQ}\mW_Q^\T, and likewise for keys and values. In self-attention, the same X\mX 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 Xq\mX_q, Xk\mX_k, and Xv\mX_v.

Listing 20.3 Multi-head attention backward pass
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 d×dd \times d matrices WQ\mW_Q, WK\mW_K, WV\mW_V, and WO\mW_O, standard multi-head attention has

4d2(20.4)4d^2\tag{20.4}

parameters, independent of the number of heads as long as the total width dd stays fixed. For self-attention on a batch of BB sequences of length TT, the dominant multiply-add count scales as

4BTd2+2BT2d.(20.5)4BTd^2 + 2BT^2d .\tag{20.5}

The first term is the four dense projections. The second term is score computation and weighted value aggregation across all heads, because Hdh=dH d_h=d. More heads change the layout and the subspaces, not this leading-order total.

Listing 20.4 Parameter and FLOP counts
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 T2T^2 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 B×H×T×TB \times H \times T \times T 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.

Key equations
Q=XqWQ,K=XkWK,V=XvWV\mQ=\mX_q\mW_Q,\qquad \mK=\mX_k\mW_K,\qquad \mV=\mX_v\mW_V
Oh=softmax⁡(QhKh⊤/dh)Vh\mO_h=\softmax(\mQ_h\mK_h^\T/\sqrt{d_h})\mV_h
Y=concat⁡(O1,…,OH)WO\mY=\operatorname{concat}(\mO_1,\ldots,\mO_H)\mW_O
#parameters=4d2,FLOPs≈4BTd2+2BT2d\#\text{parameters}=4d^2,\qquad \text{FLOPs}\approx 4BTd^2+2BT^2d
Xˉself=Xˉq+Xˉk+Xˉv\bar{\mX}_{\text{self}}=\bar{\mX}_q+\bar{\mX}_k+\bar{\mX}_v

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.

  1. Draw X\mX entering WQ\mW_Q, WK\mW_K, and WV\mW_V.

  2. Split each projected width into HH blocks.

  3. Run one attention equation per block with the same causal mask.

  4. Concatenate blocks and multiply by WO\mW_O.

Misconceptions to address.

  • "More heads always means more parameters." Not if dd 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

Exercise 20.1 ★ Why heads?

Explain why splitting attention into several heads can be more expressive than one head with the same total width.

Exercise 20.2 ★★ Shape trace

Trace the shapes from (B,T,d)(B,T,d) through projection, split into HH heads, attention, concatenation, and output projection.

Exercise 20.3 ★★ Parameter and FLOP budget

Derive the 4d24d^2 parameter count and the leading self-attention cost in (20.5).

Exercise 20.4 ★★★ Gradient-check MHA

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