Multi-query and grouped-query attention
Be able to explain how MQA/GQA reduce the KV cache and what is sacrificed.
Prerequisites
- EThe KV cacherequired
Intuition
In ordinary multi-head attention (MHA) every head has its own Q, K and V. The KV cache scales with the number of heads — and becomes the memory hog at long contexts.
MQA (multi-query): all the heads share one K/V pair. The cache shrinks by the number of heads (32×), but the quality drops noticeably.
GQA (grouped-query): the heads are divided into groups that share K/V. With 32 Q heads and 8 KV groups the cache becomes 4× smaller at nearly unchanged quality.
GQA is the standard in Llama 3, Mistral and Qwen — one of the few changes that gives a large gain at almost no price.
Formal
The KV cache size:
where = the layers, = the number of KV heads (not Q heads), = the sequence length, = the batch.
An example, a 7B model (L = 32, d_head = 128, fp16, T = 8 192, B = 1):
| Variant | H_kv | The cache |
|---|---|---|
| MHA | 32 | 4.3 GB |
| GQA (4 groups) | 8 | 1.1 GB |
| MQA | 1 | 0.13 GB |
With GQA four times as many simultaneous users fit on the same card — which translates directly into the cost per user.
What is sacrificed: expressiveness. Every group has to share the same keys and values, so heads within a group cannot specialise as freely. Empirically GQA with 4–8 groups is nearly impossible to tell from MHA, while MQA shows in the quality.
Uptraining: an MHA model can be converted to GQA by averaging the K/V projections within each group and then continuing to train on ~5 % of the original compute budget (Ainslie et al. 2023).
Code
def kv_cache_gb(layers=32, kv_heads=32, d_head=128, seq=8192, batch=1, bytes_per_number=2):
return 2 * layers * kv_heads * d_head * seq * batch * bytes_per_number / 1e9
for name, h in (("MHA", 32), ("GQA-8", 8), ("GQA-4", 4), ("MQA", 1)):
print(f"{name:6s} {kv_cache_gb(kv_heads=h):6.2f} GB")
# MHA 4.29 GB
# GQA-8 1.07 GB
# GQA-4 0.54 GB
# MQA 0.13 GB
# How many simultaneous users fit on an 80 GB card after the model weights (14 GB)?
for name, h in (("MHA", 32), ("GQA-8", 8)):
print(name, int((80 - 14) / kv_cache_gb(kv_heads=h)))
# MHA 15
# GQA-8 61
Mastery means
- Explains MHA, MQA and GQA
- Computes the KV cache size for each variant
- Knows what is sacrificed
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — GQA: Training Generalized Multi-Query Transformer Models — arXiv (open access; licence per article)
- arXiv — Fast Transformer Decoding: One Write-Head is All You Need (MQA) — arXiv (open access; licence per article)