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.
1import torch2B, T, d_model, n_heads = 2, 5, 768, 123d_head = d_model // n_heads # 644q = torch.randn(B, T, d_model) # output of W_q, one block56qh = q.view(B, T, n_heads, d_head).transpose(1, 2) # (B, 12, T, 64)7back = qh.transpose(1, 2).reshape(B, T, d_model) # concat heads again8print(qh.shape, torch.equal(back, q)) # torch.Size([2, 12, 5, 64]) TrueSplitting 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
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_headis 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_modelfor queries but uses fewer K/V heads, soW_kandW_vare narrower. d_headset independently. Gemma 7B hasd_model = 3072but 16 heads of 256, so the concatenated heads are 4,096 wide andW_omaps 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_headsand keepd_model, what happens to parameters?" — Nothing;d_headhalves and the projections stayd_model × d_model. - "Does the KV cache depend on
n_heads?" — It depends onn_kv_heads × d_head, which equalsd_modelin 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.