Attention, Part 2: MQA, GQA and MLA

Every token leaves its keys and values behind in memory (the KV cache), and plain multi-head attention leaves a lot. How multi-query, grouped-query and multi-head latent attention shrink it, explained from zero, built from scratch, checked against a real model's layer, and measured: DeepSeek-V3 stores 68.6 KiB per token where plain attention would need 4,880.

Part 1 built attention from scratch and ended on its cost: memory. This part is about three ideas that cut that memory, sometimes by more than 60 times, and which ones today's models use.

As before, every technical word gets a yellow box the first time it appears.

The memory problem: the KV cache

A chatbot writes its answer one token at a time. To write each new token, it runs attention: the new token's query is compared with the keys of every earlier token, and their values are mixed.

Those earlier keys and values never change. So instead of recomputing them at every step, the model computes them once and keeps them in memory.

KV cache:kept in memory,one K and V per tokenTheVKcatVKsatVKonVKtheVKQthe new queryWriting the next word after "The cat sat on the"Blue: the newest token. Its K and V are added to the cache. Its query is compared with every stored key, then thrown away.
How the KV cache works. Every earlier token has left a key and a value in memory. The newest token adds its own, and its query is compared with every stored key.

Notice what is not in the cache: the query. A query belongs to the token being written right now. It is used once, for this one step, and thrown away. Only keys and values are kept.

That is the opening for this whole part: a model can keep many query heads while storing far fewer key and value heads.

How big is it?

With plain multi-head attention (MHA), every head in every layer stores one key and one value for every token:

KV cache per token=2×L×H×dh×b\text{KV cache per token} = 2 \times L \times H \times d_h \times b

where:

  • the 22 counts one key and one value;
  • LL is the number of layers;
  • HH is the number of key/value heads in each layer (in plain MHA, the same as the number of query heads);
  • dhd_h is the length of each head's key and value vectors;
  • bb is the bytes per number (2 for the usual 16-bit numbers).

For Llama 2 7B, 2×32×32×128×2=524,2882 \times 32 \times 32 \times 128 \times 2 = 524{,}288 bytes: 512 KiB for every single token. A 4,096-token conversation needs 2 GiB, and a GPU serving 30 conversations at once needs 60 GiB just for caches.

Why reading the cache is the real cost

When a model writes one token, the arithmetic is small, but it must read all its weights and the whole KV cache from GPU memory. Reading memory has a speed limit, called memory bandwidth, and for writing it is usually the limit that matters:

tstep  ≳  bytes of weights+bytes of KV cachememory bandwidtht_{\text{step}} \;\gtrsim\; \frac{\text{bytes of weights} + \text{bytes of KV cache}}{\text{memory bandwidth}}

where tstept_{\text{step}} is the time to write one token, and memory bandwidth is how many bytes per second the GPU can read (about 2,000 to 3,000 GB per second on a modern data-centre GPU).

Every byte the cache saves is a byte that does not need to be read, at every step, for every conversation on that GPU. That is why the three ideas below matter so much.

The cache, more than the arithmetic, limits how many people a GPU can serve and how long their conversations can be. (The LLM inference series goes deeper.) Look at the formula again: the only part we can shrink without changing the model's size much is HH, the number of key/value heads, or what each head stores. That is exactly what the three ideas do.

MHAown K, V per head8 K/V heads cachedGQAK VK Vgroups share K, V2 K/V heads cachedMQAK Vall share one K, V1 K/V head cachedMLAlatent cK, V rebuilt from cc + RoPE key cachedorange: query heads (computed fresh, never cached). Aqua and blue: what each token leaves in the KV cache.
Four ways to organise the heads. Orange dots are query heads: they are computed fresh and never stored. Aqua boxes are key/value heads: these are what each token leaves in the cache. MLA stores one compressed vector (blue) instead.

Multi-query attention (MQA): one key/value head for everyone

In 2019, Noam Shazeer proposed the most extreme answer: keep all the query heads, but give the whole layer just one key head and one value head, shared by every query head.

In the formula, HH becomes 11:

MQA cache per token=2×L×1×dh×b\text{MQA cache per token} = 2 \times L \times 1 \times d_h \times b

So the cache shrinks by a factor equal to the number of heads.

Falcon-7B is a real example: 71 query heads sharing one key/value head. It stores 8 KiB per token, against 512 KiB for Llama 2 7B.

The price is quality. Every query head must find what it needs in the very same keys and values, and models trained this way tend to be measurably worse. The DeepSeek-V2 paper says it plainly: MQA and GQA "require a smaller magnitude of KV cache, but their performance does not match MHA."

The original paper measured both sides of that trade, on English-to-German translation:

Grouped-query attention (GQA): the compromise everyone adopted

In 2023, Ainslie et al. proposed a middle ground: split the query heads into groups, and give each group its own key/value head.

GQA cache per token=2×L×Hkv×dh×b,group size g=HqHkv\text{GQA cache per token} = 2 \times L \times H_{kv} \times d_h \times b, \qquad \text{group size } g = \frac{H_q}{H_{kv}}

where HqH_q is the number of query heads and HkvH_{kv} the number of key/value heads.

GQA is a dial between the two extremes:

  • With Hkv=HqH_{kv} = H_q (every query head has its own), GQA is plain multi-head attention.
  • With Hkv=1H_{kv} = 1 (one for everyone), GQA is multi-query attention.

The paper showed that GQA gets "quality close to multi-head attention with comparable speed to MQA". It also showed that an existing MHA model can be converted to GQA: average the key and value heads inside each group, then train briefly (about 5% of the original training) to recover.

As an equation, converting one group of gg key heads into a single shared key head is just an average of their weight matrices:

Wkgroup=1g∑i∈groupWk(i),Wvgroup=1g∑i∈groupWv(i)W_k^{\text{group}} = \frac{1}{g} \sum_{i \in \text{group}} W_k^{(i)}, \qquad W_v^{\text{group}} = \frac{1}{g} \sum_{i \in \text{group}} W_v^{(i)}

where Wk(i)W_k^{(i)} and Wv(i)W_v^{(i)} are the key and value matrices of head ii in the original model.

It became the default. Llama 3 (8 key/value heads), Qwen2.5 (Qwen2.5-0.5B: 14 query heads, 2 key/value heads), Gemma 3 (27B: 32 query heads, 16 key/value heads) and Mistral all use it.

GQA from scratch

One way to build GQA is to copy each key/value head once for every query head in its group, then run normal attention. That works, but making those copies wastes the memory we just saved. The better way is to group the query heads instead, so the small cache is read as it is:

python
def gqa(q, k, v, causal=True):
    """q: (B, Hq, T, d); k, v: (B, Hkv, T, d) with Hq a multiple of Hkv.
    Each group of Hq/Hkv query heads shares one key/value head. No copy of K or V is made."""
    B, Hq, T, d = q.shape
    Hkv = k.shape[1]
    g = Hq // Hkv
    qg = q.view(B, Hkv, g, T, d)                                       # group the query heads
    scores = torch.einsum('bhgqd,bhkd->bhgqk', qg, k) / math.sqrt(d)
    if causal:
        mask = torch.triu(torch.ones(T, T, dtype=torch.bool, device=q.device), 1)
        scores = scores.masked_fill(mask, float('-inf'))
    out = torch.einsum('bhgqk,bhkd->bhgqd', torch.softmax(scores, -1), v)
    return out.reshape(B, Hq, T, d)

The shapes read like this: B is how many texts at once, T the number of tokens, d the head length, and Hq and Hkv the numbers of query and key/value heads.

Proof against PyTorch. PyTorch's built-in attention function needs the copying approach, so I gave it copied keys and values and compared it with the function above, for all three cases:

plain text
MHA (8 KV heads)   vs PyTorch with repeated K/V: max |diff| = 4.4e-16
GQA (2 KV heads)   vs PyTorch with repeated K/V: max |diff| = 4.4e-16
MQA (1 KV head)    vs PyTorch with repeated K/V: max |diff| = 6.7e-16

Differences around 10−1610^{-16} are the smallest rounding errors a computer can make here: same answers.

Proof on a real model

A toy test only proves the toy. So I opened Qwen2.5-0.5B, a real model that uses GQA with 14 query heads sharing 2 key/value heads. I captured the input to its first attention layer and recomputed that layer by hand, from the model's own weights:

  1. the query, key and value matrices (including their small added constants, called biases);
  2. the rotary position embedding, RoPE (explained below);
  3. the grouped attention function above;
  4. the output matrix.

Then I compared my result with what the model itself produced:

plain text
Qwen2.5-0.5B layer 0 (14 query heads share 2 KV heads): our GQA vs the model's own output,
max |diff| = 6.6e-07 (outputs are up to 0.5)

6.6×10−76.6 \times 10^{-7} is the rounding noise of 32-bit numbers. So the real layer computes exactly what those 13 lines compute.

Does a smaller cache make writing faster?

At every step, the model reads the whole KV cache. Fewer key/value heads means less to read, so each step should be faster. I measured one attention step for 8 conversations at once, each with 4,096 tokens of history and 32 query heads, on an Apple M5 Pro GPU:

0 ms1 ms2 ms3 ms32 KV heads (MHA): 3.08 ms, reads 537 MB of cache3.08 ms32 KV heads (MHA)reads 537 MB8 KV heads (GQA): 1.85 ms, reads 134 MB of cache1.85 ms8 KV heads (GQA)reads 134 MB1 KV head (MQA): 1.24 ms, reads 17 MB of cache1.24 ms1 KV head (MQA)reads 17 MB
Time for one attention step over a 4,096-token cache. Fewer key/value heads means less to read and a faster step, but not in proportion.
Key/value headsCache read per stepTime per attention step
32 (MHA)537 MB3.08 ms
8 (GQA)134 MB1.85 ms
1 (MQA)17 MB1.24 ms

Faster, clearly, but not 4 or 32 times faster. Each step has a fixed cost (starting the work on the GPU) that does not shrink, and once the cache is small that fixed cost dominates.

The bigger win is memory: a cache 4 times smaller fits 4 times more conversations, or 4 times more history, on the same GPU.

Multi-head latent attention (MLA): compress instead of share

GQA saves memory by throwing information away: query heads in the same group are forced to use identical keys and values. DeepSeek-V2 (2024) tried something different: what if each token stored a compressed version of its keys and values, from which every head's own keys and values can be rebuilt?

xtokenlatent cvia W_dkvRoPE keyvia W_kr + RoPEthe whole KV cachekeys, head 1..hvalues, head 1..hrebuilt by W_uk and W_uvTraining: rebuild everyhead's K and V from c.Inference: never rebuild.Fold W_uk into the query andW_uv into the output, andattend directly over c.Same answer, measured:max |diff| = 3.9e-16
Multi-head latent attention. Each token is compressed into a small latent vector c, plus a small key that carries position (RoPE). Only those two are cached. Every head's keys and values can be rebuilt from c.

The MLA equations

For each token with vector xx, MLA computes two small things and caches only them:

c=x WdkvkR=RoPE⁡(x Wkr)c = x\,W_{dkv} \qquad\qquad k^{R} = \operatorname{RoPE}(x\,W_{kr})

where:

  • cc is the latent: WdkvW_{dkv} ("down-projection for keys and values") shrinks the token to dcd_c numbers (512 in DeepSeek-V3);
  • kRk^R is a small position key of drd_r numbers (64 in DeepSeek-V3), shared by all heads.

When a head hh needs its keys and values, it rebuilds them from the latent with two up-projections:

kh=c Wuk(h)vh=c Wuv(h)k_h = c\,W_{uk}^{(h)} \qquad\qquad v_h = c\,W_{uv}^{(h)}

Here are the same three equations in the paper, numbered (9), (10) and (11):

So the cache per token is just:

MLA cache per token=L×(dc+dr)×b\text{MLA cache per token} = L \times (d_c + d_r) \times b

For DeepSeek-V3: 61×(512+64)×2=70,27261 \times (512 + 64) \times 2 = 70{,}272 bytes, or 68.6 KiB.

(DeepSeek also compresses the queries through a small latent, to save memory during training. Queries are never cached, so that part does not change the KV cache, and my implementation leaves it out.)

The trick that makes it fast: never rebuild at all

Rebuilding every head's keys and values at every step would cost a lot of extra work. MLA avoids it with a small piece of algebra called absorption.

A head's attention score is its query dotted with a key. Put in the rebuilt key and rearrange:

qh⋅kh  =  qh⋅(c Wuk(h))  =  (qh Wuk(h)⊤)⋅cq_h \cdot k_h \;=\; q_h \cdot \big(c\,W_{uk}^{(h)}\big) \;=\; \big(q_h\,W_{uk}^{(h)\top}\big) \cdot c

The left side needs every cached token's key to be rebuilt. The right side multiplies the one new query by Wuk(h)⊤W_{uk}^{(h)\top} once, then compares it with the cached latents directly. Same number, much less work.

The value side works the same way. The head's output is a weighted sum of values, with weights aja_j from softmax:

∑jaj vh,j  =  ∑jaj cjWuv(h)  =  (∑jaj cj) Wuv(h)\sum_j a_j\, v_{h,j} \;=\; \sum_j a_j\, c_j W_{uv}^{(h)} \;=\; \Big(\sum_j a_j\, c_j\Big)\, W_{uv}^{(h)}

So the weighted sum is taken over the small latents, and Wuv(h)W_{uv}^{(h)} is applied just once at the end. The full keys and values are never built.

Here is that inference path from my implementation:

python
def absorbed(self, x):
    """Inference path: cache only c (d_c) and the RoPE key (d_r) per token; never rebuild K or V."""
    T = x.shape[0]; pos = torch.arange(T)
    c, kr = self.W_dkv(x), rope(self.W_kr(x), pos)                     # <- the entire KV cache
    W_uk = self.W_uk.weight.view(self.h, self.dn, self.dc)
    W_uv = self.W_uv.weight.view(self.h, self.dv, self.dc)
    q_nope = self.W_q(x).view(T, self.h, self.dn)
    q_lat = torch.einsum('qhd,hdc->qhc', q_nope, W_uk)                  # absorb W_uk into the query
    q_rope = torch.stack([rope(t, pos) for t in self.W_qr(x).view(T, self.h, self.dr).unbind(1)], 1)
    scores = (torch.einsum('qhc,kc->hqk', q_lat, c) + torch.einsum('qhd,kd->hqk', q_rope, kr)) / math.sqrt(self.dn + self.dr)
    mask = torch.triu(torch.ones(T, T, dtype=torch.bool), 1)
    A = torch.softmax(scores.masked_fill(mask, float('-inf')), -1)
    o_lat = torch.einsum('hqk,kc->qhc', A, c)                            # attend in latent space
    o = torch.einsum('qhc,hvc->qhv', o_lat, W_uv)                        # then up-project once
    return self.W_o(o.reshape(T, -1))

Proof. I compared it with the slow path that rebuilds every key and value, using the same weights:

plain text
MLA: full path vs compressed-cache path, max |diff| = 3.9e-16
numbers cached per token per layer: full K and V 640, latent + RoPE key 80 (8.0x smaller)

Identical answers, with an eighth of the cache in this small example.

Why position needs its own key

You may wonder why MLA has that separate little position key. The reason is RoPE.

position 0turned by 0°position 1turned by 30°position 2turned by 60°position 3turned by 90°Dashed: the vector before RoPE. Blue: after RoPE. Each position turns it a little more (here 30° per position).
RoPE in a picture. The same vector is turned a little more at each position. Comparing two turned vectors tells the model how far apart they are.

For one pair of numbers (x1,x2)(x_1, x_2) inside a vector at position mm, RoPE does:

(x1′x2′)=(cos⁡mθ−sin⁡mθsin⁡mθcos⁡mθ)(x1x2)\begin{pmatrix} x_1' \\ x_2' \end{pmatrix} = \begin{pmatrix} \cos m\theta & -\sin m\theta \\ \sin m\theta & \cos m\theta \end{pmatrix} \begin{pmatrix} x_1 \\ x_2 \end{pmatrix}

where mθm\theta is the angle: θ\theta is a fixed small angle (different for each pair of numbers), and mm is the position.

Here is the problem. If RoPE were applied to the rebuilt keys kh=c Wuk(h)k_h = c\,W_{uk}^{(h)}, a rotation would sit between Wuk(h)W_{uk}^{(h)} and the query, and that rotation is different for every position. Then there is no single matrix to move over to the query side, and the absorption trick breaks.

I tested exactly that mistake: apply RoPE to the rebuilt keys and the queries, then compare with the compressed path:

plain text
with RoPE applied to the reconstructed keys, the compressed path is wrong by up to 0.03
(outputs up to 0.53)

An error of about 6%, from one misplaced rotation.

The DeepSeek-V2 paper found exactly this problem and states it in one sentence:

RoPE itself comes from the RoFormer paper, which explains it with this picture:

DeepSeek's fix is to split the job. The latent carries the content and is never rotated. A small separate key carries the position, and only it gets RoPE. The score simply adds the two parts:

scoreh,j=(qhWuk(h)⊤)⋅cj  +  qhR⋅kjRdn+dr\text{score}_{h,j} = \frac{\big(q_h W_{uk}^{(h)\top}\big)\cdot c_j \;+\; q^R_h \cdot k^R_j}{\sqrt{d_n + d_r}}

where qhRq^R_h is the head's small position query, kjRk^R_j is token jj's cached position key, and dnd_n and drd_r are the lengths of the content and position parts. (This is exactly the scores = ... line in the code above.)

How big is the difference?

Here are four real models, using each one's published settings and 2 bytes per number:

Llama 2 7B (MHA, 32 KV heads): 512.0 KiB per tokenLlama 2 7BMHA, 32 KV heads512.0 KiBLlama 3.1 8B (GQA, 8 KV heads): 128.0 KiB per tokenLlama 3.1 8BGQA, 8 KV heads128.0 KiBDeepSeek-V3 (671B) (MLA, latent 512 + RoPE 64): 68.6 KiB per tokenDeepSeek-V3 (671B)MLA, latent 512 + RoPE 6468.6 KiBFalcon 7B (MQA, 1 KV head): 8.0 KiB per tokenFalcon 7BMQA, 1 KV head8.0 KiB
KV cache per token for four real models. Llama 2 7B uses plain multi-head attention; Llama 3.1 8B uses GQA; DeepSeek-V3 uses MLA; Falcon 7B uses MQA.
ModelAttentionFormulaCached per token
Llama 2 7BMHA2 × 32 layers × 32 heads × 128 × 2 B512 KiB
Llama 3.1 8BGQA2 × 32 layers × 8 heads × 128 × 2 B128 KiB
DeepSeek-V3 (671B)MLA61 layers × (512 + 64) × 2 B68.6 KiB
Falcon 7BMQA2 × 32 layers × 1 head × 64 × 2 B8 KiB

The DeepSeek-V3 row is the striking one. It is a 671-billion-parameter model with 128 attention heads, yet it stores less per token than an 8-billion-parameter Llama. With plain multi-head attention and the same heads, it would need 4,880 KiB per token, about 71 times more.

The DeepSeek-V2 paper says MLA's cache is as small as GQA with only 2.25 groups, and reports that MLA, unlike GQA and MQA, "achieves better performance than MHA". That claim comes from their own tests at large scale. MLA is also harder to build and serve, which is part of why GQA is still the most common choice.

Where does 2.25 come from? DeepSeek-V2 sets the latent to dc=4dhd_c = 4 d_h and the RoPE key to dhR=dh/2d_h^R = d_h / 2, where dhd_h is one head's size. Set MLA's cache equal to GQA's and solve for the number of groups ngn_g:

(dc+dhR) l=(4+12)dh l=92 dh l=2 ng dh l⟹ng=94=2.25(d_c + d_h^R)\, l = \left(4 + \tfrac{1}{2}\right) d_h\, l = \tfrac{9}{2}\, d_h\, l = 2\, n_g\, d_h\, l \qquad\Longrightarrow\qquad n_g = \tfrac{9}{4} = 2.25

where ll is the number of layers and the factor 2 on the GQA side counts keys and values.

Which one to use?

Cache per tokenQualityHow hardUsed by
MHAlargestthe baselinesimplestolder models (Llama 2 7B, GPT-2)
MQAsmallestnoticeably worsesimpleFalcon, PaLM
GQAin between, a dialclose to MHAsimpleLlama 3, Qwen2.5, Gemma 3, Mistral
MLAsmallas good as MHA or better, per DeepSeekharder: latent, absorption, separate RoPE keyDeepSeek-V2/V3/R1, Kimi K2

All four keep one thing the same: every token still looks at every earlier token. They shrink what each token stores, not how many tokens are looked at. Cutting that is the subject of Part 3.

The impact, and where you meet it

These three ideas changed how every large model is built and served:

  • GQA is the default. Llama 2's largest models adopted it, and Llama 3, Qwen2.5, Gemma 3 and Mistral followed. If you run an open model today, it almost certainly uses GQA.
  • MLA made very large models cheap to serve. DeepSeek-V2 and V3 used it to serve very large mixture-of-experts models with a small cache, and Moonshot's Kimi models adopted it.
  • MQA was used by Falcon, PaLM and the original small Gemma 2B, where speed and memory mattered most.

Use cases in one line each:

  • Chat services: a 4 to 8 times smaller cache (GQA) means 4 to 8 times more simultaneous users per GPU, or much longer conversations.
  • Long documents and code: at 128K tokens the cache, not the model, fills the GPU; MLA and GQA make such contexts affordable.
  • On-device models: phones and laptops have little memory, so small KV caches (GQA or MQA) decide how long a conversation can be.
  • Reasoning models: models that "think" for thousands of tokens before answering write very long outputs, and every one of those tokens reads the cache.

Summary

  • When a model writes, it keeps every earlier token's keys and values in the KV cache. Queries are used once and never stored.
  • Plain MHA stores 2×L×H×dh×b2 \times L \times H \times d_h \times b bytes per token: 512 KiB for Llama 2 7B.
  • MQA shares one key/value head across all query heads (Falcon 7B: 8 KiB per token) but loses quality. GQA shares in groups, a dial between MHA and MQA.
  • Our GQA matched PyTorch to 10−1610^{-16} and reproduced Qwen2.5-0.5B's real first layer to 6.6×10−76.6 \times 10^{-7}.
  • Fewer key/value heads made each step faster (3.08, 1.85, 1.24 ms) and, more importantly, the cache smaller.
  • MLA stores a small compressed latent plus a small RoPE key, and absorbs the up-projections into the query and output so keys and values are never rebuilt. The fast path matched the slow path to 3.9×10−163.9 \times 10^{-16}.
  • RoPE must stay out of the latent: putting it on the rebuilt keys broke the shortcut by about 6%.
  • DeepSeek-V3 stores 68.6 KiB per token; plain MHA would need 4,880 KiB.
  • Writing is limited by memory reads, tstep≳bytes read/bandwidtht_{\text{step}} \gtrsim \text{bytes read} / \text{bandwidth}, which is why shrinking the cache speeds it up (MQA's decoder became about 12× faster in its paper).
  • In the papers: MQA lost 0.2 BLEU for a 12× faster decoder; GQA-XXL matched MHA-XXL (47.1 vs 47.2) more than 5 times faster; DeepSeek-V2 cut its KV cache by 93.3%.
Run it yourself

Everything above comes from code/attention/part2_kv_variants.py, which also contains the full MLA class with both the slow (rebuild) path and the fast (absorbed) path.

bash
pip install torch transformers
python part2_kv_variants.py     # writes results/part2.json

References

  1. N. Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019.
  2. J. Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP 2023.
  3. DeepSeek-AI. DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model. 2024.
  4. J. Su et al. RoFormer: Enhanced Transformer with Rotary Position Embedding. 2021.
  5. Model configurations: Llama 2 7B, Llama 3.1 8B, DeepSeek-V3, Falcon 7B, Qwen2.5-0.5B.