56 GQA / MQA
Context
When you serve an LLM, most of the memory goes on the KV cache (see #49), and it is proportional to the number of heads for which you store keys and values. Full multi-head attention (#32) keeps separate K, V for each of the h heads — expensive in memory on long context.
The idea and the mechanism
The observation: you need many query heads (they ask different "questions"), but the keys and values can be shared. MQA takes that to the limit: one shared K/V head for all queries — the cache drops by a factor of h, but the model loses expressiveness and trains worse and less stably. GQA is the compromise: split the h query heads into g groups, and let each group share one K/V head. At g = 1 this is MQA, at g = h it is ordinary MHA. A bonus: GQA can be cheaply "uptrained" from an existing MHA checkpoint rather than trained from scratch.
linear algebra · systems Where the KV-cache saving comes from
The size of the KV cache per request is proportional to the number of KV heads \( n_{kv} \) (keys and values per layer, per KV head, per token):
Varying only \( n_{kv} \) gives you the whole spectrum — from full MHA to MQA at the extreme:
For example, 32 query heads and 8 KV groups → every four queries share one K/V → 4× less KV cache at nearly MHA quality. The number of query heads (and hence the cost of the attention matmul itself) does not change — the saving is specifically in the cache memory, that is, in how many requests fit into a batch.
PyTorch GQA: "replicating" g KV heads up to h for the matmul
import torch
# q: [B, h, T, d]; k,v: [B, g, T, d] (g KV groups, g divides h)
def gqa(q, k, v, h, g):
rep = h // g # how many queries per KV head
k = k.repeat_interleave(rep, dim=1) # [B, h, T, d] — one shared K per group
v = v.repeat_interleave(rep, dim=1)
a = (q @ k.transpose(-1,-2)) / q.size(-1)**0.5
return a.softmax(-1) @ v # store only g KV heads, compute as if h
Why it matters
GQA made long context cheap in memory and became the de facto standard in nearly every modern open LLM (LLaMA-2 70B, LLaMA-3, Mistral). It is part of the same fight over the KV cache as vLLM (#49), prefix caching (#54) and MLA (#53) — only here the cache is squeezed by reducing the number of KV heads.
Connections
GQA/MQA is a direct modification of the multi-head attention from #32: the same mechanics, but with K/V heads shared between queries. What changes is not the idea of attention but its memory "bookkeeping" at inference.
Both attack the cost of attention, but from different sides: FlashAttention cuts the memory traffic while computing attention, GQA cuts the size of the KV cache between steps. In production they are combined.
DeepSeek's MLA is another way to squeeze the same cache: GQA reduces the number of KV heads (by sharing them), while MLA compresses each KV into a low-rank latent. Different axes for shrinking one bottleneck.
Questions worth asking
Why does sharing K/V barely hurt quality, while sharing queries does?
Query heads ask different questions of the context — their diversity carries a lot of information, and cutting it is costly. Keys and values are the shared "content", and empirically there is far more redundancy in it: several queries can look at the same K/V without much loss. So it is the KV heads that get reduced.
If MQA (g=1) is so compact, why bother with GQA at all?
MQA is too aggressive: one K/V head for everybody loses expressiveness, gives a noticeable drop in quality and trains unstably (especially when up-training large models). GQA keeps several groups — enough to hold quality nearly at the MHA level while the cache is still several times smaller. g is the "memory ↔ quality" trade-off knob.
Can you get a GQA model without training from scratch?
Yes — that is the paper's key practical contribution. An existing MHA checkpoint is "converted": the K/V heads within each group are averaged into one, and then it is uptrained — briefly fine-tuned (a small fraction of the original budget). That way almost any existing MHA model can be moved to GQA cheaply, without paying for a full training run.
What to read in the original
Read selectively: from GQA (arXiv 2305.13245) — the idea of groups and of uptraining from an MHA checkpoint, plus the "quality vs number of groups" curve; from MQA (Shazeer, 2019) — the original claim that one write-head is enough, and why the cache was already a bottleneck back then.