Fast Attention and Long Context
Grouped-query attention
Grouped-query attention gives every small group of attention heads one shared set of keys and values, which is why almost every model released since 2023 uses it.
- 11 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Grouped-query attention gives each small group of heads one shared set of notes.
The classroom you sat in
Forty students, one subject, and a stack of textbooks. Buying forty books is expensive and most of them sit closed.
Buying one book for the whole class is too far the other way. Nobody can work.
So every bench of four shares a book. Cheap enough to afford, and close enough that everyone can read. Grouped-query attention is that bench.
Where this sits
The previous lesson, multi-query attention, shared one set of keys and values across every head. Keys and values are the model's stored notes about each earlier word.
It worked, and it hurt. Quality slipped, and training got twitchy at large scale.
Grouped-query attention keeps the same idea and turns the dial back. Thirty-two heads, eight shared sets, four heads per set.
MHA (32 sets) GQA (8 sets) MQA (1 set)
------------- ------------ -----------
head 1 -> set 1 head 1 -┐ head 1 -┐
head 2 -> set 2 head 2 -┼-> set 1 head 2 -┤
head 3 -> set 3 head 3 -┤ head 3 -┼-> the
head 4 -> set 4 head 4 -┘ head 4 -┤ only
... ... ... -┤ set
head 32 -> set 32 head 32 -┐ set 8 head 32 -┘
biggest notes middle smallest notes
best quality nearly as good quality slipsWhy the middle wins
Four times fewer notes than the full version. Almost all of the speed of the extreme version. Quality that measures close to the full version.
That is a rare shape for a trade-off. Most compromises give you half of each benefit. This one gives you most of both, which is why it became the default within a year.
The other reason it won
There is a practical reason too, and it is worth knowing.
Large models are split across several graphics cards. Each card holds some of the heads. With one shared set of notes, every card needs its own copy, so the saving mostly evaporates.
With eight sets and eight cards, each card gets one set and keeps the full saving. The number eight was picked to match how machines are wired, not only for quality.
A useful thing to check
You can read the choice straight off any open model. Its settings file lists the number of attention heads and the number of key-value heads.
If they are equal, the model uses the original scheme. If the second is smaller, it is grouped. If it is one, it is the extreme version.
Where you have already seen this
- Llama, Mistral, Qwen and most other open models you can download today.
- A long chat that keeps replying quickly instead of slowing down.
- A 7-billion-parameter model that fits a long document on a modest card.
Remember this
- Grouped-query attention shares one set of notes per small group of heads.
- It keeps nearly the quality of the original and nearly the speed of the extreme version.
- A group count of eight also matches how models are split across cards.
What to learn next
- Multi-head latent attention — compressing the cache rather than sharing it.
- Sliding-window attention — bounding the cache with a fixed window.
- Memory-bound vs compute-bound — why cache size sets decoding speed.
Developer — Code and libraries.
Setup
pip install torchWritten against PyTorch 2.5.1, Python 3.10. Runs on CPU.
The dial, and the kernel that knows about it
import torch, torch.nn.functional as F
torch.manual_seed(0)
B, H, N, D = 1, 8, 32, 16
def run(n_kv):
q = torch.randn(B, H, N, D)
k = torch.randn(B, n_kv, N, D)
v = torch.randn(B, n_kv, N, D)
# the explicit way: copy each kv head to the query heads in its group
g = H // n_kv
manual = F.scaled_dot_product_attention(
q, k.repeat_interleave(g, dim=1), v.repeat_interleave(g, dim=1), is_causal=True)
return q, k, v, manual
for n_kv in (8, 4, 2, 1):
q, k, v, manual = run(n_kv)
kv_numel = k.numel() + v.numel()
name = {8: "MHA ", 1: "MQA "}.get(n_kv, "GQA ")
print(f"{name} kv_heads={n_kv} group size={H//n_kv} "
f"cache elements={kv_numel:>6} relative={kv_numel/(2*B*H*N*D):.3f}")
# PyTorch can do the grouping inside the kernel, with no copies
q = torch.randn(B, H, N, D)
k = torch.randn(B, 2, N, D)
v = torch.randn(B, 2, N, D)
builtin = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True)
manual = F.scaled_dot_product_attention(q, k.repeat_interleave(4, 1),
v.repeat_interleave(4, 1), is_causal=True)
print()
print("enable_gqa=True matches the manual expansion:", torch.allclose(builtin, manual, atol=1e-6))
print("max difference:", (builtin - manual).abs().max().item())
print()
print("what real models chose (from each config.json on the Hub)")
rows = [("Llama-3.1-8B", 32, 8), ("Mistral-7B-v0.1", 32, 8), ("Mixtral-8x7B-v0.1", 32, 8),
("Qwen3-30B-A3B", 32, 4), ("gpt-oss-20b", 64, 8), ("DeepSeek-V3", 128, 128)]
print(f"{'model':<20}{'q heads':>9}{'kv heads':>10}{'group':>7}")
for name, qh, kvh in rows:
print(f"{name:<20}{qh:>9}{kvh:>10}{qh//kvh:>7}")MHA kv_heads=8 group size=1 cache elements= 8192 relative=1.000 GQA kv_heads=4 group size=2 cache elements= 4096 relative=0.500 GQA kv_heads=2 group size=4 cache elements= 2048 relative=0.250 MQA kv_heads=1 group size=8 cache elements= 1024 relative=0.125 enable_gqa=True matches the manual expansion: True max difference: 4.172325134277344e-07 what real models chose (from each config.json on the Hub) model q heads kv heads group Llama-3.1-8B 32 8 4 Mistral-7B-v0.1 32 8 4 Mixtral-8x7B-v0.1 32 8 4 Qwen3-30B-A3B 32 4 8 gpt-oss-20b 64 8 8 DeepSeek-V3 128 128 1
Reading the output
The cache column is a straight division. Halving the key/value heads halves the cache. There is nothing subtle in the memory arithmetic; the subtlety is all in the quality cost.
enable_gqa=True gives the same numbers without the copies. The difference of 4e-07 is float32 reassociation. Use the flag: repeat_interleave allocates a tensor group times larger, in the hot path, for no reason. The flag is documented as experimental and, as of PyTorch 2.13, works with the flash and math kernels on CUDA tensors.
The last table is the interesting one. Four of the six models chose 8 key/value heads. Qwen3-30B-A3B went to 4. And DeepSeek-V3 shows 128 and 128, which looks like plain multi-head attention.
It is not. DeepSeek-V3 keeps full per-head keys and values at compute time but caches a compressed form instead, which the config file cannot express. That is multi-head latent attention, the next lesson.
Reading it off any model yourself
from transformers import AutoConfig
c = AutoConfig.from_pretrained("meta-llama/Llama-3.1-8B")
print(c.num_attention_heads, c.num_key_value_heads) # 32 8 -> GQA, group of 4This needs network access and, for gated repositories, a Hugging Face login. The two field names are stable across model families in transformers, which is one of the more useful conventions in the library.
Converting an existing model
You cannot drop key/value heads and expect the model to survive. The published recipe has two steps.
First, mean-pool the key and value projection matrices within each group into one matrix. Averaging preserves far more than selecting one member of the group.
Second, uptrain: continue pre-training on the original data mixture. The paper uses 5% of the original pre-training compute, which is small next to a fresh run and large next to a fine-tune.
Skipping the second step gives a model that produces text and fails benchmarks.
Common mistakes
Choosing a group count that does not divide the head count. 32 % 5 != 0, and the kernel will refuse. PyTorch requires num_heads_q % num_heads_kv == 0.
Choosing key/value heads below the tensor-parallel degree. With n_kv = 4 on 8 GPUs, the cache must be replicated across pairs of devices and you lose part of the saving. Match or exceed the parallel degree.
Benchmarking GQA on prefill. Prefill is compute-bound; the win is in decode. Measure tokens per second during generation at a realistic context length, not prompt processing.
Assuming the cache is the only saving. Fewer key/value heads also means fewer projection parameters and slightly less work in the projections. It is a small effect and not why anyone does this.
Comparing across models by group ratio alone. gpt-oss-20b has a group of 8 but a head dimension of 64, and half its layers use a 128-token window. Cache size is a product of four numbers, not one ratio.
Try it yourself
Add n_kv = 3 to the loop and read the error PyTorch gives you. Then set H = 12 and confirm n_kv = 3 now works with a group size of 4. Finally, time enable_gqa=True against repeat_interleave at N = 4096 on a GPU and look at peak memory rather than wall clock.
What to learn next
- Multi-head latent attention — compressing the cache rather than sharing it.
- Sliding-window attention — bounding the cache with a fixed window.
- Memory-bound vs compute-bound — why cache size sets decoding speed.
Researcher — Mathematics and papers.
Definition
Partition $h$ query heads into $g$ groups of size $h/g$. Group $j$ shares one key projection $W^K_j$ and one value projection $W^V_j$. For head $i$ in group $j$:
$$ O_i = \operatorname{softmax}!\left(\frac{Q_i K_j^{\top}}{\sqrt{d_h}}\right) V_j $$
$g = h$ recovers multi-head attention; $g = 1$ recovers multi-query attention. GQA is the one-parameter family between them, from Ainslie, Lee-Thorp, de Jong, Zemlyanskiy, Lebrón and Sanghai (2023), GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints (arXiv:2305.13245, EMNLP 2023).
The abstract states the two contributions precisely: "a recipe for uptraining existing multi-head language model checkpoints into models with MQA using 5% of original pre-training compute", and GQA as "a generalization of multi-query attention which uses an intermediate ... number of key-value heads", with the finding that "uptrained GQA achieves quality close to multi-head attention with comparable speed to MQA".
Uptraining
Conversion mean-pools the projection matrices within each group:
$$ W^K_j = \frac{g}{h} \sum_{i \in \text{group } j} W^K_i $$
and likewise for $W^V$. Mean-pooling beats both selecting a single head and random initialisation, which is the empirical result that makes conversion practical at all. A short continuation of pre-training then recovers the remaining gap.
The consequence for practitioners is that GQA is not a decision you must make before spending millions on pre-training. It is a decision you can make afterwards, for 5% more.
Why the cost curve is so favourable
Cache bytes scale as $\Theta(g)$ while measured quality degrades slowly until $g$ becomes very small. The paper's interpolation shows most of the MQA speed-up is already obtained at $g = 8$, with quality close to MHA. Two mechanisms are usually offered.
Attention heads are known to be redundant. Michel, Levy and Neubig (2019), Are Sixteen Heads Really Better than One? (arXiv:1905.10650) prune most heads at test time with modest loss. GQA shares rather than prunes, which is gentler.
Query diversity is preserved in full. Every head keeps its own $W^Q_i$, so heads can still ask different questions of a shared key space. The lost capacity is in what can be stored, not in what can be asked.
Systems fit
With tensor parallelism of degree $t$, heads are sharded across devices. Per-device KV cache is
$$ \text{bytes/device} = 2 b L \frac{\max(g, t)}{t} d_h S B \cdot \frac{t}{\max(g,t)} = 2 b L \, \frac{g}{t} \, d_h S B \quad \text{when } g \ge t $$
and stops improving once $g < t$, because the shared heads must be replicated. The prevalence of $g = 8$ across Llama, Mistral, Mixtral and gpt-oss is not a coincidence: it is the largest common tensor-parallel degree on an 8-GPU node.
Where the frontier moved
GQA reduces the constant in front of $S$. It does not change the fact that cache grows linearly with context. Three later directions attack that:
- Latent compression. MLA caches a low-rank latent vector instead of per-head keys and values, decoupling cache size from head count entirely.
- Windowing. Sliding-window attention bounds the cache by a constant on most layers.
- Cache quantisation and eviction. Orthogonal to all of the above, and typically the cheapest remaining win in a deployed system.
GQA remains the default because it is free at inference, free to adopt from an existing checkpoint, and costs almost nothing in quality. It is one of the few architectural choices in modern language models with essentially no dissenting camp.
What to learn next
- Multi-head latent attention — compressing the cache rather than sharing it.
- Sliding-window attention — bounding the cache with a fixed window.
- Memory-bound vs compute-bound — why cache size sets decoding speed.