Only keys and values are ever cached, so sharing them eight query heads to one shrinks the cache eightfold for a quality loss that is hard to measure.
Why you'd careThe thing you have already noticed
You loaded a 70B model in bf16 across two 80 GB cards, saw the weights take 140 GB, and found the server would admit only a handful of long conversations before reporting no free KV blocks. Or you did the cache arithmetic for a classic multi-head model and got a number larger than the weights themselves. That is the problem grouped-query attention solves. Every token in context holds a key and a value for every layer and every KV head, for the life of the conversation. Llama 3 70B, with 8 KV heads, spends about 320 KB per token. At 64 heads it would be 2.5 MB, and a single 8,000-token conversation would need 20 GB of cache on its own.
In and outWhat goes in, what comes out
| In | The layer's hidden state plus its projection matrices. Q projects to H heads, 64 in Llama 3 70B; K and V project to G heads, 8. G must divide H. MQA is G = 1, classic MHA is G = H. |
|---|---|
| Process | Each KV group is broadcast to the H/G query heads assigned to it — logically a repeat, and in a well-written kernel not an actual memory copy. Attention then runs per query head exactly as it would under MHA. |
| Out | The same [S, d_model] attention output MHA would produce, and a cache of shape [2, layers, positions, G, head_dim], smaller by a factor of H/G. Nothing downstream can tell which variant produced it. |
Queries stay fully independent; what is shared is the index the queries search over. Heads inside a group still ask different questions, but they are all answered from one key-value store, so they can no longer specialise in what to store. That capacity was given up at training time and is not recoverable at serving time — the checkpoint was trained with the grouping already in place.
ConceptThe idea underneath
The cost of a cached token is exact and worth memorising: bytes_per_token = 2 * n_layers * G * head_dim * dtype_bytes. The 2 is for K and V. Nothing in that expression mentions the number of query heads, which is the whole insight — queries are recomputed every step and never stored, so they are free as far as memory is concerned. Keys and values are stored for the life of the sequence.
Multi-query attention took the extreme position: one shared key/value head for all queries. It cut the cache by the full head count and cost measurable quality, because a single shared key space has to serve every head's different retrieval need. Grouped-query attention interpolates. Partition H query heads into G groups and give each group its own K and V. G = H is MHA, G = 1 is MQA, and G = 8 sits almost exactly on the knee of the curve: eight times less cache, quality within noise of MHA on standard benchmarks.
Existing MHA checkpoints do not need retraining. Mean-pool the key and value projections within each group to initialise the shared heads, then continue pretraining for a small fraction of the original compute — around 5 percent in the GQA paper. That cheapness is why GQA spread across the whole open ecosystem within about a year instead of waiting for a new model generation.
One deployment constraint bites. Tensor parallelism splits heads across GPUs, so G must be at least the TP degree. Serve a G = 8 model on 16-way tensor parallelism and the KV heads get replicated across pairs of GPUs: you pay twice the cache and gain nothing.
At a glanceSee it
Query heads outnumber key and value heads, and only the key/value side is ever cached.
The knobsHyperparameters and nuance
- num_key_value_headsG itself. 8 across every Llama 3 size and in Mistral 7B; Qwen 2.5 is not uniform, using 8 only from 14B up, 4 at 7B and 2 on the small models. Lower means less cache and faster decode with gently worsening simultaneous retrieval; 1 costs quality you can measure on multi-hop tasks.
- tensor_parallel_sizemust not exceed G, or KV heads replicate across GPUs. This is the commonest silent waste in multi-GPU serving, because nothing errors and nothing logs it.
- kv_cache_dtypesetting it to fp8 halves the cache again on top of GQA and stacks multiplicatively. It costs a little precision in long-context retrieval and needs per-tensor scaling factors to be correct.
- head_dimappears linearly in the cache formula and is increasingly set independently of hidden size, so it is now a direct memory knob rather than a derived one.
- group size, num_attention_heads divided by num_key_value_headshow many query heads compete for one key space. 4 at the 7B/8B scale (Llama 3 8B and Mistral 7B are both 32/8) and 8 at the 70B scale (Llama 3 70B is 64/8); Qwen 2.5 7B is an odd 28/4 = 7. The number to check first when a model is oddly bad at tracking several things at once.
EffectHow this stage moves the answer
Reducing G degrades gently and specifically. What goes first is not fluency or reasoning but simultaneous distinct lookups: keeping four variables straight while writing code that uses all of them, comparing two tables in the same document, tracking who said what across a long multi-party transcript. With fewer key spaces, heads in a group compete for the same index and the model resolves the competition by blending. From outside that looks like conflation — the right structure with two facts swapped, or an attribute from row three applied to row seven. At G = 8 this is essentially undetectable against MHA. At G = 1 it is measurable, which is why MQA lost to GQA in production. If your workload leans on many simultaneous distinct retrievals, group size is worth checking before concluding the model is simply weaker.
EvalsWhat it does to your measurements
The metric this stage moves is not a quality score, it is concurrency. Cache bytes per token determine how many sequences fit alongside the weights, and that sets the throughput ceiling of every load test you run. Measure it directly from the formula times your p95 context length, against free memory after weights are loaded. On quality, the GQA paper's result is the one to trust: uptrained GQA models land within noise of MHA on summarization and question answering while MQA shows a consistent small deficit. The evaluation trap is context length. A suite that runs at 2,000 tokens finds no difference between G = 1 and G = 8 and also no memory pressure, because at that length nothing is constrained. Both effects only exist at the lengths you actually serve.
Failure modesWhen it goes wrong
- Free KV blocks exhausted long before the GPU is compute-saturatedcache per token times concurrency exceeds what remains after weights. The fix is G, KV dtype or max context, not a faster card.
- Moving from 8-way to 16-way tensor parallelism yields no extra usable contextG is 8, TP is 16, so key/value heads are replicated across GPU pairs.
- Memory spikes inside the attention kernel despite a small cachea naive implementation materialising the repeated K and V up to H heads before calling the kernel, undoing the entire saving.
- A converted checkpoint is fluent but noticeably worse at retrievalMHA projections mean-pooled into groups without the uptraining step that repairs them.
- Long-context accuracy drops after enabling fp8 KV cachequantization error in a shared key space affects every query head in that group at once, so there is less averaging to hide it.
PapersWhere this comes from
- Fast Transformer Decoding: One Write-Head is All You NeedShazeer, 2019. arXiv:1911.02150. Introduced multi-query attention and established that decode speed is set by KV memory bandwidth, which is the premise this whole stage rests on.
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head CheckpointsAinslie et al., 2023. arXiv:2305.13245. Defined grouped-query attention and showed existing MHA checkpoints can be uptrained into it for roughly 5 percent of pretraining compute, which is why it spread so fast.
- DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language ModelDeepSeek-AI, 2024. Introduced multi-head latent attention, which compresses keys and values into a shared low-rank latent instead of sharing heads — the main production alternative to GQA and a different point on the same tradeoff.