Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Explain Grouped-Query Attention (GQA) vs Multi-Head Attention—memory and quality differences?
What you need to know
Three layouts of the same idea
MHA: n_q = 32, n_kv = 32 → cache ∝ 32 headsGQA: n_q = 32, n_kv = 8 → cache ∝ 8 heads (4× smaller)MQA: n_q = 32, n_kv = 1 → cache ∝ 1 head (32× smaller)MQA came from Shazeer (2019), "Fast Transformer Decoding: One Write-Head is All You Need". GQA came from Ainslie et al. (2023), who also showed you can convert an existing MHA checkpoint — mean-pool the K/V heads within each group, then "uptrain" with about 5% of the original pretraining compute — and recover quality close to MHA while keeping most of MQA's speed.
What changes and what does not
- KV cache — shrinks by
n_q / n_kv. This is the main win. - Decode speed — decode is memory-bound, and at long context the cache can be a large part of what each step reads, so a smaller cache means faster steps and bigger batches.
- Attention FLOPs — unchanged. Each query head still computes a full row of scores against keys.
- Parameters —
W_kandW_vshrink. In Llama-3-8B each is4096 × 1024instead of4096 × 4096, saving about 25M parameters per layer: small next to the cache saving.
How it looks in code
1import torch2B, T, n_q, n_kv, hd = 1, 6, 32, 8, 1283q = torch.randn(B, n_q, T, hd)4k = torch.randn(B, n_kv, T, hd) # only 8 K heads are cached5v = torch.randn(B, n_kv, T, hd)67group = n_q // n_kv # 4 query heads per K/V head8k_exp = k.repeat_interleave(group, dim=1) # (1, 32, 6, 128) for the math9v_exp = v.repeat_interleave(group, dim=1)10out = torch.nn.functional.scaled_dot_product_attention(q, k_exp, v_exp, is_causal=True)11print(out.shape, k.numel() / k_exp.numel()) # torch.Size([1, 32, 6, 128]) 0.25Only the 8-head k and v are stored. Query heads 0–3 read KV head 0, heads 4–7 read KV head 1, and so on. Recent PyTorch versions accept enable_gqa=True in scaled_dot_product_attention, so optimised kernels do this sharing without materialising the expanded copy.
Where the field went next
GQA with 8 KV heads became the default for open models from Llama 2 70B onward. DeepSeek-V2 (2024) introduced multi-head latent attention (MLA), which caches one small compressed vector per token and reconstructs K and V from it; the paper reports a 93.3% smaller KV cache than their earlier 67B dense model. Some models also share one KV cache across neighbouring layers.
A real-life example
A team serves a chatbot on a 70B model across 8 H100s (640 GB). The fp16 weights take about 140 GB, leaving roughly 500 GB for cache. Llama-3-70B has 80 layers and head_dim 128.
MHA (64 KV heads): 2 × 80 × 64 × 128 × 2 bytes = 2.5 MiB per token → 20 GiB per 8k chatGQA (8 KV heads): 320 KiB per token → 2.5 GiB per 8k chatWith MHA the node holds about 25 full-length conversations; with GQA, about 200. For a chatbot with a million daily users, that 8× difference is the difference between a cluster of nodes and a building's worth of them.
Follow-up questions to expect
- "Does GQA reduce attention FLOPs?" — Hardly. Every query head still scores every key; only the K/V projections and the cache get smaller.
- "How do you choose the number of KV groups?" — It is a quality-versus-memory knob; 8 is common because it also divides evenly across 8-way tensor parallelism.
- "Why is MQA worse?" — All heads must read the same keys and values, which limits how differently heads can attend, and the original GQA paper also reports MQA being less stable to train.