Transformer Architecture Q&A

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_o maps 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):

Text
MHA(x) = concat(head_1, ..., head_n) · W_o       = Σ_i  head_i · W_o^(i)
Python
import numpy as nprng = np.random.default_rng(0)T, n_heads, d_head = 4, 3, 2d_model = n_heads * d_headheads = [rng.normal(size=(T, d_head)) for _ in range(n_heads)]W_o = rng.normal(size=(d_model, d_model))concat_then_project = np.concatenate(heads, axis=1) @ W_oblocks = np.split(W_o, n_heads, axis=0)          # 3 blocks of shape (2, 6)sum_of_heads = sum(h @ Wb for h, Wb in zip(heads, blocks))print(np.allclose(concat_then_project, sum_of_heads))   # True

The 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):

  1. 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.
  2. Compute locally — each GPU computes its heads and multiplies by its own W_o block, giving a partial d_model-wide output.
  3. 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_o and 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_o makes the write location learnable.
  • "Does W_o have a bias?" — In GPT-2 yes; Llama-style models drop all linear biases.
  • "Why is W_o row-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.