10. GQA: Grouped-Query Attention (2023, Ainslie et al.)
arXiv: 2305.13245
Core Problem
MQA is too extreme—all heads share 1 KV group, and quality drops noticeably. But MHA's KV cache is too large. Is there a middle ground? For example, split 96 heads into 8 groups, each sharing 1 KV set—reducing the cache while preserving some head diversity?Method
GQA's core idea is an intermediate state: not 1 KV group (MQA), not n_heads KV groups (MHA), but G KV groups (where 1 < G < n_heads).How it works:
- Split the n_heads query heads into G groups
- Query heads within each group share the same K and V
- Groups remain independent
- KV cache reduced to 1/G (e.g., 87.5% reduction when G=8)
- Uptraining requires only 5% of the original pretraining compute
- Quality "close to multi-head attention with comparable speed to MQA"
For example, LLaMA-2 70B: n_heads=64, n_kv_heads=8 (G=8). Each group of 8 query heads shares 1 KV set, reducing the KV cache to 1/8.
Even better: the paper shows an existing MHA model can be uptrained into GQA using only ~5% of the original pretraining compute!
Key Numbers
Impact
GQA is the "golden midpoint" between MHA and MQA. Mainstream models like LLaMA-2/3, Gemma, and Mistral all use GQA. It enables large models to maintain quality while substantially speeding up inference, making it a key technique for industrial deployment. The uptraining method also lets existing models "upgrade" to GQA without training from scratch.Takeaway
> GQA's way of thinking: "don't pick one of two—find a third path." MHA and MQA are two extremes—one has a huge cache, the other loses quality. GQA asked a key question: if diversity comes from "each group independent" rather than "each head independent," how many groups are needed to reach the quality sweet spot? The answer: 8–12 groups suffice. It's like an orchestra—you don't need 96 soloists; 8 sections are enough. The insight: optimization problems often aren't at the extremes of parameter space, but at some intermediate balance point.---
arXiv: 2305.13245