Era 6 · Generative models and systems · 2019 / 2023

56 GQA / MQA

Fast Transformer Decoding: One Write-Head is All You Need (MQA · Shazeer, 2019) · GQA: Training Generalized Multi-Query Transformer Models… (Ainslie et al. · Google, 2023)
🟧 read selectively~1 horiginal ↗
The gist in 20 seconds. The KV cache grows with the number of attention heads — and that is the bottleneck of inference. MQA (2019): all query heads share ONE K/V head → a tiny cache, but a drop in quality and unstable training. GQA (2023): the happy medium — h query heads are split into g groups, and each group shares a K/V head. g is a tunable knob (1 = MQA, h = MHA), typically giving ~4× less KV at quality close to full MHA. Standard in LLaMA-2/3 and Mistral.

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):

\[ \text{KV cache} \;\propto\; 2\,L\,n_{kv}\,d_{h}\,T \]

Varying only \( n_{kv} \) gives you the whole spectrum — from full MHA to MQA at the extreme:

\[ \text{MHA: } n_{kv}=h \qquad \text{GQA: } n_{kv}=g \qquad \text{MQA: } n_{kv}=1 \quad\Rightarrow\quad \text{cache} \downarrow\ \tfrac{h}{g}\times \]

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
MHA (n_kv = h) QKV GQA (n_kv = g) ↑2 groups MQA (n_kv = 1) 1 for all fewer KV heads → smaller cache → more requests per batch query heads (blue) stay the same; only K/V are shared (count below)
The number of query heads is constant; what shrinks is the number of KV heads they share: MHA — one each, GQA — one per group, MQA — one for all. Fewer KV heads = less cache.
Analogy. Translators (query heads) and reference books (K/V). In MHA every translator has their own full reference set — precise, but bulky. In MQA there is one shared set for everybody — compact, but there is a queue at the desk and details get lost. GQA gives one set per small group of translators: almost as precise, but with several times fewer copies — and cheaper to carry around (to keep in memory).

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

← simplifies32. Transformer

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.

↔ a different axis48. FlashAttention

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.

↔ compresses even harder53. DeepSeek V3 / R1

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.