How Grouped-Query Attention Shrinks the KV Cache
Grouped-query attention shares key/value heads across query heads to cut the KV cache. How GQA sits between MHA and MQA, with the memory math and code.
In an earlier post on the KV cache, I listed grouped-query attention as the first lever to reach for when GPU memory runs out. Then I moved on in a single line: “pick a GQA model.” That is the right advice and a bad explanation. It skips what the lever actually is, why it works without wrecking quality, and why nearly every open model shipped in the last two years already has it turned on.
This post fills that gap. It is for engineers who serve or fine-tune LLMs and want to know what “GQA” in a model card is buying them, where the memory goes, and what the tradeoff against plain multi-head attention really costs. There is a little linear algebra and a short PyTorch snippet, but the core idea is a counting argument you can do on a napkin.
The one term in the cache formula that GQA attacks
Recall the per-token size of the KV cache (the keys and values stored for every past token) from the last post:
bytes/token = 2 x num_layers x num_kv_heads x head_dim x dtype_bytesFour of those five terms are hard to touch at serving time:
num_layersandhead_dimare fixed by the architecture.dtype_bytesyou can cut with quantization.- The leading
2is just keys plus values.
That leaves num_kv_heads, the number of distinct key/value heads the model keeps per layer. GQA is the design choice that makes this number small without making the model dumb. Because the term sits inside a product, cutting it in half halves the whole cache.
To see why that number can be small, you have to look at how attention wires queries to keys and values in the first place.
Three ways to wire keys and values: MHA, MQA, and GQA
Multi-head attention (MHA). Standard multi-head attention, from the original transformer, gives every attention head its own three projections: a query, a key, and a value. Twelve heads means twelve independent key heads and twelve value heads, and at decode time (while generating tokens) you cache all of them. Each head learns to attend to a different pattern, which is where multi-head attention gets its expressiveness. It is also where the cache gets its size.
Multi-query attention (MQA). In 2019 Noam Shazeer noticed that the queries are the expensive-to-share part, not the keys and values. His multi-query attention keeps all the query heads but collapses the key and value heads down to a single shared pair, so every query head attends against the same keys and values. The cache shrinks by the full head count, and decoding gets noticeably faster because there is far less memory to move per step. The problem is that one shared KV head is a real bottleneck. MQA can lose accuracy, and models trained with it can be less stable, so you often could not just swap it into an existing model.
Grouped-query attention (GQA). GQA is the compromise, and the reason it feels obvious in hindsight. Instead of one KV head for everyone (MQA) or one per query head (MHA), you split the query heads into groups and give each group its own shared KV head. Eight query heads with two KV heads, for example, means two groups of four, each group sharing a pair. You get to pick where you sit on the line between the two extremes.
The middle panel carries the argument. Quality tracks the number of KV heads gently, but memory tracks it steeply. In practice, going from eight KV heads to two barely moves accuracy, while it cuts the cache by four. That asymmetry is the entire case for GQA.
How much GQA saves on real models
The savings are exactly the ratio of query heads to KV heads, because that ratio is how many query heads now share each cached KV head. Take a real model. Llama-2-70B has 64 query heads, so under plain multi-head attention it would cache 64 key heads and 64 value heads per layer. It ships with 8 KV heads instead. Its cache is 8x smaller than the multi-head version would be, at essentially the same benchmark scores.
The pattern repeats across the models you are likely to run. Both Llama 3 8B and Mistral 7B use 32 query heads with 8 KV heads, a 4x reduction. Falcon-40B and most other current open models made the same call. When a model card lists num_attention_heads: 32 and num_key_value_heads: 8, that 4:1 gap is GQA. It tells you the KV cache is a quarter of what the head count alone would suggest.
The 8x on the 70B model is not a rounding detail. Plug it back into the cache arithmetic, and it is the difference between a batch size of one and a batch size of eight on the same card, or between a 4k context and a 32k one. Memory you do not spend on redundant KV heads is memory you can spend on longer context or more concurrent requests.
How the sharing looks in code
At inference time, the cache stores only the small number of KV heads. To run attention, each of those KV heads has to be matched against several query heads, so the implementation expands the cached KV heads back up to the query-head count right before the attention math. In the Llama code this is a helper usually called repeat_kv, which repeats each KV head n_rep times:
import torch
def repeat_kv(kv: torch.Tensor, n_rep: int) -> torch.Tensor:
"""Expand cached KV heads to line up with the query heads.
kv shape: (batch, n_kv_heads, seq_len, head_dim)
n_rep = n_query_heads // n_kv_heads
returns: (batch, n_kv_heads * n_rep, seq_len, head_dim)
"""
b, n_kv, seq, head_dim = kv.shape
if n_rep == 1: # this is plain multi-head attention
return kv
return (
kv[:, :, None, :, :]
.expand(b, n_kv, n_rep, seq, head_dim)
.reshape(b, n_kv * n_rep, seq, head_dim)
)The detail that matters is expand. It creates a view, not a copy, so the repeated heads do not each take their own memory. Only the n_kv real heads live in the cache; the expansion is a broadcast that happens per step and is then thrown away.
That is why GQA is a genuine memory win and not just bookkeeping. What you pay to store is n_kv_heads, what you compute against is n_query_heads, and the two are decoupled. Set n_rep to the query-head count and you have MQA; set it to 1 and you are back to multi-head.
Converting existing models: uptraining
An obvious worry is that GQA sounds like something you must decide before pretraining. That would make it useless for the pile of MHA checkpoints already in the wild. The 2023 paper that named and popularised GQA, Ainslie et al. from Google Research, spends most of its length on exactly this problem, and the fix is neat:
- Take a trained multi-head model and mean-pool its key and value heads within each intended group. The eight KV heads that will become one group, for example, are averaged into a single head to initialise it.
- Continue training for a short while so the model adjusts to the shared heads.
They call this uptraining, and it costs about 5% of the original pretraining compute (paper). The result lands where you want it: quality close to the original multi-head model, at decoding speed close to multi-query. That recipe is why GQA spread so fast. Nobody had to retrain from scratch to get it.
Tradeoffs and failure modes
GQA is close to free, but “close to free” is not “free.” A few things are worth knowing before you treat it as a dial.
- It is baked in, not a serving flag. Unlike KV quantization or paging the cache, you cannot switch GQA on for a model that was not built or uptrained with it. It is a reason to pick a GQA model up front, not a lever you pull in production.
- Group count is a quality knob, and the low end costs you. More KV heads is closer to multi-head quality, and fewer is closer to multi-query. The 8:1 ratio on the 70B is a well-tested sweet spot, but pushing toward one KV head brings back the accuracy and stability problems that made plain MQA unattractive. If you are training or uptraining your own model, do not assume more compression is free.
- The savings are on KV heads only. GQA does nothing about the other multipliers. A long context or a deep agent transcript still grows the cache linearly, and a GQA model with a 100k-token conversation can still run you out of memory. It lowers the slope, not the shape.
- It composes, so measure the stack, not the piece. GQA multiplies cleanly with int8 KV quantization and with paging, and the wins stack. That is good, but it means a memory budget has to account for all three together, rather than crediting any one of them with all the headroom.
How GQA fits into sizing a deployment
When I size a deployment, GQA is not really a decision anymore; it is a filter. I would not pick a memory-tight serving target on a model that lacks it, the same way I would not pick one I could not quantize. The interesting choices come after that filter, and that is where an application engineer actually has room to move. GQA sets the slope of the cache. How much context you feed and how many requests you batch set where you land on that slope.
That downstream part is where I spend most of my time. On Archi, the RAG copilot I built for CMS computing operations, a single answer can pull thousands of tokens of logs and tickets into the window. The cache math is a daily constraint there, and GQA is the reason a long-context answer fits on the GPU I have rather than the one I would have to ask for. The same arithmetic shows up on the HPC and GPU side of the CMS workflow tooling I maintained. GQA is a good example of an architectural choice that looks like a footnote in a model card and turns out to decide whether your batch size is one or eight.
Diagrams by M. Hassan Ahmed, released under CC0. No external image was used in this post; the figures are original work by the author.