One dial from MHA to MQA
The whole family is a single continuum. Multi-head attention (MHA) gives every one of the h query heads its own key and value head, so the cache stores h K/V pairs per token per layer. Multi-query attention (MQA) collapses that to a single shared K/V pair for all heads. GQA sits between them: keep all h query heads, but partition them into g groups that each share one K/V pair.
So g = h is MHA, g = 1 is MQA, and any divisor of h in between is a valid GQA design. Crucially the query side never changes — queries are recomputed each step and thrown away, so they are cheap. The only thing g moves is how many distinct K and V vectors accumulate in the cache. That single fact is why choosing g is almost entirely a systems decision about memory and bandwidth, with only a light touch of modelling quality on the other side of the scale.
Why fewer K, V heads is the right lever
Why attack the number of K/V heads specifically, and not, say, the head dimension? Because autoregressive decode is memory-bandwidth-bound, and the KV cache is the part of the byte traffic that grows without bound. Each decode step must stream two things out of memory: the model weights (once) and the entire KV cache (once). The arithmetic per token is tiny by comparison, so time-per-token tracks bytes moved, not FLOPs.
The useful way to see this is arithmetic intensity — FLOPs performed per byte read. Roofline analysis says throughput is bandwidth-limited whenever intensity sits below the hardware’s FLOP-to-bandwidth ratio, which decode almost always does. Shrinking the KV cache by h/g cuts the bytes each step reads, which raises arithmetic intensity and pushes decode toward the compute-bound regime where the accelerator is actually busy. Reducing d_head instead would shrink model capacity everywhere; reducing K/V heads targets exactly the bandwidth bottleneck and leaves the query-side expressivity intact. That surgical fit is why GQA, not some other trim, became the standard.
What the smaller cache buys: a throughput example
The cache saving translates directly into concurrency, which is where serving economics live. Take a 70B-class model: L = 80 layers, d_head = 128, bf16 (p = 2 bytes). Per token per layer a single K/V head costs 2 · d_head · p = 512 bytes, so one sequence at S = 8192 tokens holds:
per-seq KV = 2 · L · n_kv · d_head · S · p
MHA (n_kv = 64): 2·80·64·128·8192·2 ≈ 20.0 GiB / sequence
GQA (n_kv = 8): 2·80· 8·128·8192·2 ≈ 2.5 GiB / sequenceNow suppose that after loading weights you have a 40 GiB budget left for KV cache on the accelerator. Under MHA you fit just 40 / 20 = 2 concurrent 8k-token sequences; under GQA with g = 8 you fit 40 / 2.5 = 16. That is an eightfold jump in batch size from the same memory — and because decode throughput scales with how many sequences you can run in parallel before hitting the memory wall, it is roughly an eightfold jump in tokens-per-second of serving capacity too. The cache dial is really a throughput dial.
The same arithmetic explains long context from the other direction: if you fix the batch size instead of the memory, GQA lets each sequence carry roughly h/g times more tokens of history within the budget. So one choice of g simultaneously widens how many users you serve and how much context each of them gets — you spend the saving on whichever axis your product needs more.
Choosing g on the quality-versus-cache Pareto
If bigger g means less sharing and more quality, and smaller g means more saving, the sweet spot is wherever the two curves cross usefully. The empirical finding from Ainslie et al. is that the trade is pleasantly asymmetric: moving from MHA down to a modest number of groups recovers almost all of MQA’s memory and speed win while giving up close to nothing in accuracy, whereas the last step down to g = 1 (full MQA) is where quality and training stability actually start to suffer.
In Pareto terms, g = 8 sits at the knee of the frontier — you have already banked most of the cache reduction, and paying more (smaller g) buys diminishing memory returns at rising quality cost. That is why eight has become the de-facto default rather than one. The rule of thumb: pick the smallest g that keeps evaluation metrics within noise of your MHA baseline, and do not chase the last factor of two in cache size — it is the expensive half of the curve.
How g interacts with tensor parallelism
There is a second, systems-level constraint on g that is easy to miss until deployment. Large models are served with tensor parallelism (TP): the attention heads are sharded across several GPUs. Query heads divide cleanly, but the n_kv = g K/V heads have to be distributed too, and the clean case is one (or an integer number of) K/V heads per GPU.
When g is at least the TP degree and divisible by it, each GPU owns whole K/V heads and no coordination is needed. When g is smaller than the TP degree — say g = 1 (MQA) on 8 GPUs — the shared K/V must be replicated across ranks, which costs a little memory and complicates the kernel. A group count of eight is convenient precisely because it maps one K/V head per GPU on the common 8-way TP setup. This is a happy alignment, not the reason eight was chosen — the quality/cache knee came first — but it does mean you should sanity-check g against your intended TP layout before committing to a number.