Below the API

MQA, GQA, and MLA

Expert Advanced 1h 15m Difficulty 4/5

Prerequisites III.08, V.05, V.06


1. The problem they solve#

Standard multi-head attention (MHA) stores K and V for EVERY head.

  KV per token = 2 × L × n_heads × head_dim × bytes

Llama-2-70B (MHA, L=80, h=64, d_head=128, FP16):
  2 × 80 × 64 × 128 × 2 = 2,621,440 bytes = 2.5 MiB per token

At 4k context: 10.2 GB PER SEQUENCE.
On 8×A100 (640 GB) with 140 GB of weights: ~49 concurrent sequences.

The KV cache, not the weights, is the capacity limit — and it’s determined by an architectural choice made at training time.


2. The three answers#

MHA (multi-head attention)
  h query heads, h key/value heads.
  Every query head has its own K and V.

MQA (multi-query attention, Shazeer 2019)
  h query heads, 1 key/value head. ALL query heads share one K,V.
  → KV shrinks by h (e.g. 64x)
  → some quality loss

GQA (grouped-query attention, Ainslie et al. 2023)
  h query heads, g key/value heads where 1 < g < h.
  Query heads are grouped; each group shares a K,V head.
  → KV shrinks by h/g (typically 4-8x)
  → negligible quality loss at g = h/8
  → THE PRODUCTION STANDARD since 2023

MLA (multi-head latent attention, DeepSeek-V2 2024)
  Compress K and V into a shared low-rank LATENT vector.
  Cache the latent; reconstruct K and V per head on the fly.
  → KV shrinks by 10-20x
  → quality reportedly equal to or better than MHA

3. The picture#

MHA (h=8):
  queries:  q0 q1 q2 q3 q4 q5 q6 q7
  keys:     k0 k1 k2 k3 k4 k5 k6 k7     ← 8 KV heads cached
  values:   v0 v1 v2 v3 v4 v5 v6 v7

GQA (h=8, g=2):
  queries:  q0 q1 q2 q3 | q4 q5 q6 q7
             \  |  |  /    \  |  |  /
  keys:         k0            k1        ← 2 KV heads cached
  values:       v0            v1

MQA (h=8, g=1):
  queries:  q0 q1 q2 q3 q4 q5 q6 q7
             \  \  \  \ /  /  /  /
  keys:            k0                   ← 1 KV head cached
  values:          v0

MLA:
  cached:   c  (a single latent vector, dim d_c ≈ 512)
                  ↓ per-head up-projection at read time
  keys:     k0 k1 k2 k3 k4 k5 k6 k7     ← reconstructed, not cached
  values:   v0 v1 v2 v3 v4 v5 v6 v7

4. The second benefit of GQA (usually missed)#

GQA doesn’t just shrink the cache. It raises decode attention’s arithmetic intensity.

Decode attention, per step, per sequence:
  FLOPs ≈ 2 × h × S × d_head × 2      (scores + attn@V, over h query heads)
  Bytes ≈ 2 × g × S × d_head × bytes  (read K and V for g KV heads)

  intensity = FLOPs/Bytes ≈ h/g / bytes_per_elem × 2

MHA (g=h):   intensity ≈ 1
GQA (g=h/8): intensity ≈ 8
MQA (g=1):   intensity ≈ h  (e.g. 64)

GQA-8 makes decode attention 8x less memory-bound. Combined with the 8x smaller cache, that’s two independent wins from one architectural change — which is why adoption was universal and rapid.


5. MLA, mechanically#

DeepSeek’s approach. Worth understanding because it’s the most aggressive KV reduction in production.

STANDARD ATTENTION
  k_i = x W_k^(i)      cache k_i for every head i     → cache h × d_head
  v_i = x W_v^(i)      cache v_i for every head i

MLA
  c   = x W_dkv                    a shared LATENT, dim d_c (e.g. 512)
  cache ONLY c                                        → cache d_c
  
  at read time:
  k_i = c W_uk^(i)                 reconstruct per head
  v_i = c W_uv^(i)

  KEY TRICK: W_uk can be ABSORBED into W_q:
    q_i · k_i = (x W_q^(i)) · (c W_uk^(i))
              = x (W_q^(i) W_uk^(i)ᵀ) cᵀ
                  └──── precompute this ────┘
    → no per-head reconstruction needed for the scores;
      just one matmul against the cached latent.
CACHE SIZE COMPARISON (DeepSeek-V2 scale: L=60, h=128, d_head=128)
  MHA:      2 × 60 × 128 × 128 × 2 = 3.93 MiB/token
  GQA-8:    2 × 60 × 16 × 128 × 2  = 0.49 MiB/token
  MLA:      60 × (512 + 64) × 2    = 0.066 MiB/token
            (d_c=512 plus a small decoupled RoPE component)
  
  MLA vs MHA:  60x smaller
  MLA vs GQA-8: 7.4x smaller

The RoPE complication: RoPE must be applied to keys, but if keys are reconstructed from a latent, the rotation can’t be absorbed. DeepSeek’s solution is a decoupled RoPE: a small separate component of the key carries the positional information and is cached alongside the latent. That’s the +64 in the calculation above.


6. The tradeoffs, honestly#

              KV size   Quality        Compute        Complexity
MHA           1.0x      baseline       baseline       simplest
GQA-8         0.125x    ~equal         ~equal         trivial
MQA           1/h       small loss     ~equal         trivial
MLA           0.02x     equal/better*  more compute   substantial

* per DeepSeek's reported results; the architecture is theirs and
  independent replication at scale is limited.
MLA's COMPUTE COST
  reconstructing K and V (or the absorbed-matmul equivalent) costs
  extra FLOPs per decode step.
  → decode was memory-bound, so trading compute for memory is the
    RIGHT trade (Section I.07)
  → but it does mean MLA benefits less at very large batch, where
    you're closer to compute-bound

MLA's IMPLEMENTATION COST
  the absorbed-matmul formulation, the decoupled RoPE, and the
  kernel support are all nontrivial. Engine support lags.

7. What this means for you#

YOU DON'T CHOOSE THE ATTENTION VARIANT — the model does.
But you DO choose the model, and this is a first-order cost factor.

WHEN COMPARING MODELS FOR DEPLOYMENT:
  compute KV bytes per token for each (Section V.06)
  → it directly determines concurrency and therefore cost per token

  Example: two 70B models of equal benchmark quality
    Model A: MHA        → 2.5 MiB/token → ~49 concurrent @ 4k
    Model B: GQA-8      → 0.31 MiB/token → ~390 concurrent @ 4k
    → Model B costs ~8x less to serve. Same quality.

“KV bytes per token” belongs in your model evaluation criteria alongside benchmark scores. It is the single largest architecture-driven cost factor.


8. Implementation notes#

GQA IN THE KERNEL (Section IV.06)
  ✗ WRONG:  k.repeat_interleave(h // g, dim=1)
            → materializes h copies; you've thrown away the benefit
  ✓ RIGHT:  query head i reads KV head i // (h/g)
            → index arithmetic, no copy

  Any implementation doing the repeat is reading 8x more KV bandwidth
  than necessary. Check yours.

TP CONSTRAINT
  TP degree must divide g (the number of KV heads), or KV heads are
  replicated across GPUs, wasting memory.
    g=8  → TP ∈ {1,2,4,8} cleanly
    g=8, TP=16 → each KV head on 2 GPUs → 2x KV memory waste

MLA SUPPORT
  Check your engine. Support varies and the absorbed-matmul optimization
  may or may not be implemented — without it you lose much of the benefit.

9. Production implications#

  • Prefer GQA models. By 2024 essentially all new models have it; if you’re evaluating an older MHA model, factor in the 8x serving cost.
  • Compute KV bytes/token as part of model selection.
  • Verify your kernel indexes rather than repeats for GQA.
  • TP degree must divide the KV head count.
  • MLA is a genuine advance but check engine support and whether the absorbed-matmul optimization is implemented.
  • At very long context, the difference compounds: MLA at 128k context is 60x less KV than MHA, which is the difference between 3 and 180 concurrent sequences.

10. Common mistakes#

repeat_interleave for GQA. Throws away the bandwidth benefit.

TP degree not dividing the KV head count. Silent memory waste.

Comparing models on parameter count and benchmarks only. KV size is an 8-60x cost factor.

Assuming MQA’s quality loss is negligible. It’s small but real; GQA exists because MQA was slightly too lossy.

Assuming MLA is a drop-in. It requires specific kernel support.


11. Hands-on exercise#

A. Compute the difference. For five real models, compute KV bytes per token and max concurrency on a GPU you know. Rank them. Does the ranking match parameter count?

B. Verify GQA implementation. In your engine, verify that GQA is implemented by indexing rather than repeating. Measure KV bytes read per decode step and compare to the theoretical value.

C. Measure the intensity effect. For a GQA model, measure decode attention’s arithmetic intensity at g=1, g=8, g=h (simulate by configuring, or compute analytically). Confirm it scales with h/g.

D. Implement MLA. From the DeepSeek-V2 paper, implement MLA’s forward pass in PyTorch for one layer, including the absorbed matmul. Verify it matches a naive reconstruct-then-attend implementation. Measure the cache size difference.

E. TP constraint. For a model with 8 KV heads, compute the memory waste at TP=16. Is TP=16 worth it?


12. Interview questions#

  1. Explain MHA, GQA, and MQA. Quantify the KV cache difference.
  2. Why does GQA give two benefits, not one?
  3. How does MLA work, and what is the absorbed-matmul trick?
  4. Why does MLA need a decoupled RoPE component?
  5. What’s the wrong way to implement GQA in a kernel?
  6. Why must the TP degree divide the number of KV heads?
  7. Two models have equal benchmark scores; one is MHA and one is GQA-8. Which do you deploy?

13. Further reading#

  • [ESTABLISHED] Shazeer, “Fast Transformer Decoding: One Write-Head is All You Need” (2019)
  • [ESTABLISHED] Ainslie et al., “GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints” (2023)
  • [ESTABLISHED] DeepSeek-V2 and DeepSeek-V3 technical reports — MLA
  • Next: 02 — Mixture of Experts

↑↓ navigate ↵ open