Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

How are head dim, embedding dim, and number of heads related in multi-head attention?


What you need to know

The three numbers

  • d_model (embedding dimension) — the width of the residual stream; every token is a vector of this size between layers.
  • n_heads — how many separate attention patterns each layer computes.
  • d_head — the size of each head's query, key and value vectors.
Python
import torchB, T, d_model, n_heads = 2, 5, 768, 12d_head = d_model // n_heads                      # 64q = torch.randn(B, T, d_model)                   # output of W_q, one blockqh = q.view(B, T, n_heads, d_head).transpose(1, 2)   # (B, 12, T, 64)back = qh.transpose(1, 2).reshape(B, T, d_model)     # concat heads againprint(qh.shape, torch.equal(back, q))   # torch.Size([2, 12, 5, 64]) True

Splitting into heads is only a reshape. W_q produces one d_model-wide vector; the reshape reads it as 12 vectors of 64.

Why head count does not change the cost

Text
parameters of W_q, W_k, W_v, W_o = 4 × d_model²       (any n_heads)score FLOPs per layer ≈ n_heads × T² × d_head = T² × d_model   (any n_heads)

For GPT-2 small that is 4 × 768² ≈ 2.36M attention parameters per layer, whether it has 1 head or 12. What changes is the shape: more heads give more distinct attention patterns, each in a smaller subspace.

Why d_head is held near 64–128

  • If d_head is too small (say 16), each head's dot product has too few dimensions to express a precise match.
  • If there are too few heads, the layer can attend in only a few ways at once.
  • GPU kernels such as FlashAttention are tuned for head dimensions like 64, 128 and 256.

So designers pick d_head, then set n_heads = d_model / d_head. Every GPT-2 size uses 64; GPT-3 175B uses 128 (96 heads × 128 = 12,288).

Exceptions worth knowing

  • GQA keeps n_q_heads × d_head = d_model for queries but uses fewer K/V heads, so W_k and W_v are narrower.
  • d_head set independently. Gemma 7B has d_model = 3072 but 16 heads of 256, so the concatenated heads are 4,096 wide and W_o maps 4,096 back to 3,072.

A real-life example

An e-commerce site's search-ranking model is a small BERT-style cross-encoder with d_model = 384. The team compares 12 heads of 32 dimensions with 6 heads of 64. Both have the same attention parameters (4 × 384² ≈ 590K per layer) and the same FLOPs, so latency on their CPU servers is almost identical.

The only question is quality. On their offline ranking test set, they pick whichever wins, rather than assuming "more heads is better". The result is not predictable in advance, which is exactly the point to make in an interview: head count is a partition choice, not a size choice.

Follow-up questions to expect

  • "If I double n_heads and keep d_model, what happens to parameters?" — Nothing; d_head halves and the projections stay d_model × d_model.
  • "Does the KV cache depend on n_heads?" — It depends on n_kv_heads × d_head, which equals d_model in plain MHA and is smaller under GQA.
  • "Why not one big head?" — One head computes one softmax per token, so it can focus on only one mix of positions; many heads can track syntax, coreference and position at the same time.