Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Why is there an output projection (W_o) at end of multi-head attention, and what does it do?
What you need to know
What concatenation alone would give
With 12 heads of 64 dimensions, head 3's output fills dimensions 192–255 of the concatenated vector. If that vector went straight into the residual stream, head 3 could only ever write to those 64 dimensions, and it would never combine with head 7. The next layer would see 12 separate slices rather than one mixed representation.
Two jobs of W_o
- Mixing across heads. Every output dimension is a weighted sum over all heads' dimensions.
- Changing basis. Each head works in its own learned subspace (set by its
W_v).W_omaps from "head space" back to the coordinates the rest of the network reads.
The sum-of-heads view
Split W_o into n_heads row blocks of shape (d_head × d_model):
MHA(x) = concat(head_1, ..., head_n) · W_o = Σ_i head_i · W_o^(i)1import numpy as np2rng = np.random.default_rng(0)3T, n_heads, d_head = 4, 3, 24d_model = n_heads * d_head5heads = [rng.normal(size=(T, d_head)) for _ in range(n_heads)]6W_o = rng.normal(size=(d_model, d_model))78concat_then_project = np.concatenate(heads, axis=1) @ W_o9blocks = np.split(W_o, n_heads, axis=0) # 3 blocks of shape (2, 6)10sum_of_heads = sum(h @ Wb for h, Wb in zip(heads, blocks))11print(np.allclose(concat_then_project, sum_of_heads)) # TrueThe two computations are identical. This is why interpretability work treats each head as an independent unit that reads from the residual stream (through W_q, W_k, W_v) and writes to it (through its block of W_o). The pair W_v · W_o^(i) decides what a head copies and where it lands.
Its size
In Llama-3-8B, W_o is 4096 × 4096 ≈ 16.8M parameters per layer, about 537M across 32 layers — around 7% of the model. It is not a formality.
A real-life example
A company serves a chatbot to a million users on a 70B model split across 8 GPUs with tensor parallelism. The sum-of-heads view is exactly how this works in practice (the Megatron-LM scheme, Shoeybi et al., 2019):
- Split heads — each GPU holds 1/8 of the query heads and the matching K/V heads, plus the matching row block of
W_o. - Compute locally — each GPU computes its heads and multiplies by its own
W_oblock, giving a partiald_model-wide output. - All-reduce — the 8 partial outputs are summed across GPUs with one collective operation.
Because the output is a sum over heads, the GPUs never need each other's head outputs until that one sum. If W_o were applied differently, every layer would need extra communication.
Follow-up questions to expect
- "Could you remove
W_oand let the FFN mix heads?" — Partly, but then attention could only write into fixed slices of the residual stream, and every other layer would have to learn around that.W_omakes the write location learnable. - "Does
W_ohave a bias?" — In GPT-2 yes; Llama-style models drop all linear biases. - "Why is
W_orow-parallel in tensor parallelism?" — Because its input is split by head, each GPU can multiply its slice and the results only need to be summed.