08
Attention stage

Multi-Head Attention

Several attention computations run in parallel, each free to track a different pattern.

A single attention computation gives each token one distribution of weights over the sequence. Multi-head attention runs several of these in parallel, each with its own Q, K, V projections, so the model isn't limited to tracking just one kind of relationship at a time.

Attention from “it”

The
robot
carried
the
patient
because
it
was
tired
.

Three heads, three different distributions over the same sentence: one leaning toward neighboring words, one toward the coreference with "robot", one toward "patient" and "tired". Real heads aren't guaranteed to specialize this cleanly or this legibly; this is a simplified illustration of the idea, not a claim about what any specific head learns.

Combining the heads

Each head produces its own output vector for a token. Those outputs are concatenated into one long vector, then mixed by a learned matrix into a single vector of the model's normal width, so from the outside, multi-head attention still looks like a single transformation.

Per-head outputs
concat
Concatenated
× Wᴼ
Output

Every head's output is concatenated into one long vector, then multiplied by a learned matrix Wᴼ that mixes what the heads found, so multi-head attention returns to the same size as its input.

It's tempting to say a given head represents "grammar" or "coreference" or "emotion." Research has found heads that lean toward specific, describable patterns in specific models, but that's discovered after the fact by inspection, not designed in, and it doesn't generalize into a fixed role for "head 3" across every model. Treat any specific claim about what a head represents as an empirical finding about one model, not a rule.

Why one head isn't enough

A single softmax gives each token one budget of attention to divide up. One token often needs several unrelated things at once: "it" needs its referent, the verb it is the subject of, and the sentence's overall tense. One distribution can't sharpen on all of those places without spreading itself thin. Running several attention computations in parallel gives the token several budgets, each free to be spent differently.

Splitting the width instead of multiplying the cost

With hh heads and model width dd, the standard design gives every head a width of dk=dv=d/hd_k = d_v = d/h. Each head has its own projections WiQW_i^Q, WiKW_i^K, WiVW_i^V, each of shape d×d/hd \times d/h, and runs ordinary attention in that smaller space:

headi=Attention(XWiQ,  XWiK,  XWiV),i=1,…,h\text{head}_i = \text{Attention}\big(XW_i^Q,\; XW_i^K,\; XW_i^V\big), \qquad i = 1, \dots, h
MultiHead(X)=Concat(head1,…,headh) WO\text{MultiHead}(X) = \text{Concat}\big(\text{head}_1, \dots, \text{head}_h\big)\, W^O

Each head returns an n×d/hn \times d/h matrix. Concatenating them side by side gives n×dn \times d again, so the total work is about the same as one full-width head. GPT-3 has 96 heads of width 128, which add up to its model width of 12,288. Concatenation only places numbers next to each other. The mixing happens in WOW^O, a learned d×dd \times d matrix: it is the one place where what different heads found can be combined. The concatenation demo above narrows 12 numbers to 4 to keep the picture small. In the standard layout the concatenated vector is already dd wide and WOW^O is square.

Do heads specialize?

Heads are not assigned jobs. They start from different random matrices, and training pushes them apart because a varied set of heads covers more than a set of near-copies would. Pruning studies point the same way from the other side: in some trained models many heads can be removed with little loss, so not every head carries equal weight.

Heads are also where much of the memory cost of generation comes from, because each one stores keys and values for every earlier token (chapter 15). Several recent open models therefore let groups of heads share one set of keys and values, a change known as multi-query or grouped-query attention, which shrinks that cache at a small cost in quality.

The result of multi-head attention has one dd-wide row per token, and each row has now taken in information from whichever tokens its heads found relevant. The next chapter gives every token a chance to process what it gathered, on its own.

Key Takeaway

Multiple attention heads run in parallel, each free to settle into a different pattern during training. Their outputs are concatenated and mixed by a learned matrix, so the model can track more than one kind of relationship between tokens at once.