MHA, MQA, and GQA: Sharing the KV Cache
Use two head counts
Let $H_q$ be the number of query heads and $H_{kv}$ the number of key/value heads. Suppose query/key head width is $d_h$ and value width is also $d_h$ for this chapter. Ordinary MHA has $H_{kv}=H_q$. MQA has $H_{kv}=1$. GQA uses an intermediate count, commonly with $H_q$ divisible by $H_{kv}$.
For contiguous groups, query head $h$ uses K/V head $g(h)=\lfloor h/(H_q/H_{kv})\rfloor$. Its output is
Queries remain distinct, so their score rows generally differ even within a group. Sharing keys does not force all grouped queries to retrieve the same mixture.
Eight queries, two K/V heads
Take $H_q=8$ and $H_{kv}=2$. Query heads 0–3 read K/V head 0; query heads 4–7 read K/V head 1. The cache stores two key vectors and two value vectors per token per layer. It does not store eight copies of them.
An educational implementation can repeat each K/V head four times to reuse a standard MHA routine. That reproduces the function, but physically expanding and storing those repeats defeats the cache memory saving. An efficient grouped kernel reads shared data directly or broadcasts it without retaining expanded storage.
Derive the cache formula with units
where $B$ is the number of cached sequences, $L$ is layer count, $n$ is cached tokens per sequence, and $b$ is bytes per cached scalar. The factor two accounts for keys and values. Unequal key/value widths replace $2d_h$ with $d_k+d_v$; quantization scales and allocation overhead add further storage.
For one sequence, 32 layers, 8192 tokens, 32 query heads, head width 128, and two-byte cache storage, MHA uses 4 GiB. Eight K/V heads use 1 GiB. One K/V head uses 128 MiB. These are arithmetic cache estimates, not total model memory and not performance measurements.
The projection matrices also change
With residual width $d$, Q projects to $H_qd_h$ coordinates, while K and V each project to $H_{kv}d_h$. Before biases, Q/K/V parameters total $d(H_q+2H_{kv})d_h$. The output projection maps $H_qd_h$ value-output coordinates back to width $d$.
If $H_qd_h=d$, reducing K/V heads saves projection parameters as well as cache space. It does not reduce the number of query-head attention distributions or make long-context attention constant in context length. Each new query still interacts with the cached token positions it is allowed to read.
Why the cache saving can help decoding
During single-token decoding, every layer repeatedly reads historical keys and values. Sharing them can reduce memory traffic and admit larger batches or longer contexts within the same memory budget. Whether that translates into latency improvement depends on the kernel, hardware, batch size, cache layout, and other model costs.
For prefill, many queries are processed together, and matrix-multiply efficiency can dominate differently. A reduction in cache bytes is a precise structural claim; a universal speedup factor is not.
What information has been constrained?
MHA can learn separate key and value projections for every query head. GQA ties those projections within each group, reducing flexibility. The queries and output projection can adapt during training, but a trained MHA checkpoint is not generally functionally identical after averaging its K/V heads.
Fast Transformer Decoding: One Write-Head Is All You Need introduces MQA. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints studies grouped heads and an uptraining approach. Treat the initialization and subsequent training as part of conversion rather than describing it as an exact algebraic compression.
Make the group mapping explicit
# q: [Hq, nq, dh]; k,v: [Hkv, nk, dh]
group_size = Hq // Hkv
assert Hq % Hkv == 0
outputs = []
for h in range(Hq):
g = h // group_size
outputs.append(attention(q[h], k[g], v[g], allowed)[0])
The loop is a reference for correctness, not an efficient GPU implementation. The cached tensors retain their $H_{kv}$ axis throughout. Use the complexity explorer to compare the storage formulas interactively.
Try it: If 32 query heads share one K/V head, how many attention distributions are computed for one token?
Thirty-two, one for each query head. The keys and values are shared, but the queries—and therefore their scores and weights—can differ.
A different compression strategy
MLA caches a shared latent representation rather than simply reducing the number of conventional K/V heads. Its exact cache accounting needs a separate derivation.