Fast Attention and Long Context
Multi-query attention
Multi-query attention keeps many question-asking heads but gives them one shared set of keys and values, which shrinks the memory a model must re-read for every word it writes.
- 11 min read
- 3 reading levels
- Updated
Read these first
On this page 8
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Multi-query attention lets all the attention heads share one set of notes, instead of each head keeping its own.
The study group you have sat in
Eight friends prepare for an exam from one library book. Each of them could photocopy the whole chapter, so everyone carries their own bundle of paper. Eight bundles, eight times the weight, and every page is identical.
Or you keep one copy flat on the table. Everyone reads the same pages, and each person looks for a different thing in them.
The questions stay personal. The reference material is shared. That is multi-query attention.
What the heads actually keep
An attention head is one independent reader inside the model. Each head asks its own kind of question about the text so far.
To answer, a head needs two things stored for every earlier word. A key, which is how that word advertises itself. And a value, which is what that word contributes if chosen.
Ordinary attention gives every head its own keys and values. Thirty-two heads means thirty-two full sets. That store is called the KV cache — the model's saved notes on everything it has read so far.
Why the cache is the problem
To write one new word, the model reads its entire KV cache. Not part of it. All of it.
Then it writes one word and does it again for the next word. And again. The cache grows with every word, so each new word costs a little more than the last.
At long lengths the cache can outgrow the model's own weights. On an expensive graphics card, it is the thing that runs out first.
ordinary attention multi-query attention
------------------ ---------------------
head 1: keys, values head 1 ─┐
head 2: keys, values head 2 ─┤
head 3: keys, values head 3 ─┼─→ one shared
... ... ─┤ set of keys
head 32: keys, values head 32 ─┘ and values
notes stored: 32 sets notes stored: 1 setThe saving, and the cost
Sharing one set instead of thirty-two makes the notes thirty-two times smaller. Writing each word gets dramatically faster, because there is far less to re-read.
The cost is real and worth stating plainly. The heads lose the ability to store different kinds of reference material. Quality drops a little, and training can become less stable.
That trade was too blunt for most teams. The compromise that followed keeps a handful of shared sets instead of one. That is grouped-query attention, and nearly every current model uses it.
Where you have already seen this
- A chatbot that keeps a fast typing speed even after a long conversation.
- A model that fits a long document on a graphics card that "should" be too small.
- Hosted models charging far less for cached input than fresh input.
Remember this
- Every head keeps notes on every earlier word. Those notes are the KV cache.
- Multi-query attention shares one set of notes across all heads.
- It is much cheaper, slightly worse, and was replaced by a middle option.
What to learn next
- Grouped-query attention — the compromise that actually shipped.
- Multi-head latent attention — compressing the cache instead of sharing it.
- Attention — the heads and what they each learn to look for.
Developer — Code and libraries.
Setup
pip install torchWritten against PyTorch 2.5.1, Python 3.10. Everything here runs on CPU.
One module, three behaviours
The only thing that changes between multi-head, multi-query and grouped-query attention is the number of key/value heads.
import torch, torch.nn as nn, torch.nn.functional as F
torch.manual_seed(0)
d_model, n_heads, d_head, N = 256, 8, 32, 16
class Attention(nn.Module):
"""n_kv_heads = n_heads -> MHA. n_kv_heads = 1 -> MQA."""
def __init__(self, n_kv_heads):
super().__init__()
self.h, self.kv, self.dh = n_heads, n_kv_heads, d_head
self.q = nn.Linear(d_model, n_heads * d_head, bias=False)
self.k = nn.Linear(d_model, n_kv_heads * d_head, bias=False)
self.v = nn.Linear(d_model, n_kv_heads * d_head, bias=False)
self.o = nn.Linear(n_heads * d_head, d_model, bias=False)
def forward(self, x):
B, T, _ = x.shape
q = self.q(x).view(B, T, self.h, self.dh).transpose(1, 2)
k = self.k(x).view(B, T, self.kv, self.dh).transpose(1, 2)
v = self.v(x).view(B, T, self.kv, self.dh).transpose(1, 2)
cache_bytes = (k.numel() + v.numel()) * k.element_size()
# every query head reads the SAME k and v when self.kv == 1
k = k.repeat_interleave(self.h // self.kv, dim=1)
v = v.repeat_interleave(self.h // self.kv, dim=1)
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
return self.o(out.transpose(1, 2).reshape(B, T, -1)), cache_bytes
x = torch.randn(1, N, d_model)
for name, kv in (("MHA", 8), ("MQA", 1)):
m = Attention(kv)
y, cb = m(x)
params = sum(p.numel() for p in m.parameters())
kvp = m.k.weight.numel() + m.v.weight.numel()
print(f"{name}: out {tuple(y.shape)} | params {params:,} | K+V projection params {kvp:,} "
f"| cache for {N} tokens {cb:,} bytes")
print()
print("KV cache per token, in bytes, at 2 bytes per number")
print(f"{'model':<22} {'layers':>7} {'kv heads':>9} {'head dim':>9} {'B/token':>9} {'128k ctx':>10}")
specs = [ # taken from each model's config.json on the Hub
("Llama-3.1-8B (GQA)", 32, 8, 128),
("Llama-3.1-8B as MHA", 32, 32, 128),
("Mistral-7B-v0.1 (GQA)",32, 8, 128),
("gpt-oss-20b (GQA)", 24, 8, 64),
]
for name, L, kvh, dh in specs:
per_tok = 2 * L * kvh * dh * 2
print(f"{name:<22} {L:>7} {kvh:>9} {dh:>9} {per_tok:>9,} {per_tok*131072/2**30:>9.2f} GB")MHA: out (1, 16, 256) | params 262,144 | K+V projection params 131,072 | cache for 16 tokens 32,768 bytes MQA: out (1, 16, 256) | params 147,456 | K+V projection params 16,384 | cache for 16 tokens 4,096 bytes KV cache per token, in bytes, at 2 bytes per number model layers kv heads head dim B/token 128k ctx Llama-3.1-8B (GQA) 32 8 128 131,072 16.00 GB Llama-3.1-8B as MHA 32 32 128 524,288 64.00 GB Mistral-7B-v0.1 (GQA) 32 8 128 131,072 16.00 GB gpt-oss-20b (GQA) 24 8 64 49,152 6.00 GB
Reading the output
The output shape is unchanged. (1, 16, 256) in both cases. MQA is a drop-in change to the internals; nothing downstream needs to know.
The cache shrank 8x, from 32,768 to 4,096 bytes. That factor equals the head count, and it is the entire point. With 32 heads it would be 32x.
The parameter count fell too, by 114,688. That is a side effect, not the goal. Some implementations spend the saved parameters on a wider feed-forward block to keep model size constant.
repeat_interleave is a teaching device, not production code. It expands the shared keys back to eight copies in registers. A real kernel never materialises those copies; PyTorch exposes enable_gqa=True on scaled_dot_product_attention for exactly this, and vLLM has dedicated kernels.
16 GB against 64 GB at 128k context. That table is the real argument. Llama-3.1-8B holds about 16 GB of weights in bf16. Under plain multi-head attention its cache at full context would be four times its own weights, for a single sequence.
One honest caveat on that table: it counts every layer as full attention. gpt-oss-20b alternates full-attention layers with 128-token sliding-window layers, so its true cache is far below 6 GB. See sliding-window attention.
The formula worth memorising
KV cache bytes = 2 (K and V)
x bytes_per_element
x num_layers
x num_kv_heads
x head_dim
x sequence_length
x batch_sizenum_kv_heads is the only term an architect gets to shrink for free. That is the whole story of this lesson and the next three.
Common mistakes
Using num_attention_heads in the formula. The term is num_key_value_heads. On a GQA model those differ by 4x or more, and the mistake silently inflates every capacity estimate you make.
Expecting MQA to speed up prefill. Prefill is compute-bound. MQA reduces bytes, not FLOPs, so the prompt-reading phase barely moves. The gain is in decode.
Converting an MHA checkpoint by dropping heads. Keeping only head 0's key and value projections destroys quality. The published recipe mean-pools the projections across each group and then continues training briefly.
Forgetting batch size. Cache scales linearly with concurrent sequences. A serving plan that fits one 128k conversation may fit zero at batch 8.
Assuming fp16 storage. Many servers quantise the cache to fp8 or int8, halving the table again. Check what your server actually does before sizing hardware.
Try it yourself
Set n_kv_heads=2 and confirm the cache lands exactly between the two rows above. Then print k.shape before and after repeat_interleave to see where the copies appear, and where a real kernel would avoid them.
What to learn next
- Grouped-query attention — the compromise that actually shipped.
- Multi-head latent attention — compressing the cache instead of sharing it.
- Attention — the heads and what they each learn to look for.
Researcher — Mathematics and papers.
Formulation
Multi-head attention with $h$ heads computes, for head $i$:
$$ O_i = \operatorname{softmax}!\left(\frac{Q_i K_i^{\top}}{\sqrt{d_h}}\right) V_i, \qquad Q_i = X W^Q_i,\; K_i = X W^K_i,\; V_i = X W^V_i $$
$X \in \mathbb{R}^{N \times d}$ is the layer input, $d_h$ the head dimension, and $W^Q_i, W^K_i, W^V_i \in \mathbb{R}^{d \times d_h}$ the per-head projections.
Multi-query attention (Shazeer, 2019, Fast Transformer Decoding: One Write-Head is All You Need, arXiv:1911.02150) keeps $h$ distinct $W^Q_i$ but collapses the others to a single shared pair $W^K, W^V \in \mathbb{R}^{d \times d_h}$:
$$ O_i = \operatorname{softmax}!\left(\frac{Q_i K^{\top}}{\sqrt{d_h}}\right) V, \qquad K = X W^K,\; V = X W^V $$
The output projection is unchanged. Only the number of distinct key/value tensors changes, from $h$ to $1$.
Why this is a bandwidth argument, not a FLOP argument
Per decoding step, attention FLOPs are $\Theta(h \, d_h \, S)$ under both schemes, because every query head still attends over the whole context. What changes is bytes read:
$$ \text{KV bytes} = 2 \, b \, L \, n_{kv} \, d_h \, S \, B $$
with $b$ bytes per element, $L$ layers, $n_{kv}$ key/value heads, $S$ context length and $B$ batch. MQA sets $n_{kv} = 1$.
Shazeer frames the resulting quantity as memory-access-to-arithmetic ratio. For incremental decoding, standard MHA has ratio $\Theta(n/d + 1/b)$ where $n$ is sequence length and $b$ batch; MQA reduces the dominant term by a factor of $h$. The paper's own summary of the result is that models "can indeed be much faster to decode, and incur only minor quality degradation from the baseline" — read the tables in the paper for the per-configuration figures rather than trusting a single headline speed-up.
The quality cost, honestly
Shazeer reports a small but real quality loss against MHA on translation. Ainslie et al. (2023) (see grouped-query attention) confirm the gap on summarisation and question answering, and additionally report training instability at scale under MQA — loss spikes that do not occur with intermediate group counts.
The mechanistic story is that heads specialise. Different heads in a trained transformer attend to syntactic dependencies, positional offsets, repeated entities and induction patterns. Forcing one key space to serve all of them removes a degree of freedom the model demonstrably uses.
Deployment record
| Model | Year | KV heads | Note |
|---|---|---|---|
| PaLM | 2022 | 1 (MQA) | Early large-scale adoption |
| Falcon-7B | 2023 | 1 (MQA) | multi_query: true in its config |
| Falcon-40B | 2023 | 8 (GQA) | Same family, already moved to groups |
| StarCoder (GPT-BigCode) | 2023 | 1 (MQA) | Long code contexts |
| StarCoder2-7B | 2024 | 4 (GQA) | The successor switched |
| Llama 2 70B onwards | 2023 | 8 (GQA) | Where the industry settled |
Falcon is the clearest illustration. The 7B model shipped with one key/value head; the 40B model in the same release already used eight. The larger the model, the less tolerable MQA's quality cost became.
MQA was a two-year architecture. Its historical importance is that it proved the KV cache, not attention FLOPs, is the binding constraint on decoding, and it created the axis along which GQA and multi-head latent attention were later placed.
Interaction with other techniques
MQA composes with FlashAttention: one reduces cache size, the other reduces intermediate traffic, and they are orthogonal.
It composes badly with tensor parallelism. Splitting $h$ query heads across $t$ devices leaves a single key/value head that must be replicated on every device, so the per-device cache saving is $h/t$ rather than $h$. GQA with $n_{kv} = t$ is the natural fit, which is one practical reason 8 became the common group count.
Cache quantisation multiplies with all of it. An fp8 KV cache halves $b$ and is close to free in quality terms at 8 bits, though int4 caches degrade measurably on long-context retrieval.
What to learn next
- Grouped-query attention — the compromise that actually shipped.
- Multi-head latent attention — compressing the cache instead of sharing it.
- Attention — the heads and what they each learn to look for.