Skip to content
AI-grafen
FAI engineeringTransformer architecture· about 90 min· fast-moving, sources checked often· verified 2026-09-20· EN

Multi-query and grouped-query attention

Be able to explain how MQA/GQA reduce the KV cache and what is sacrificed.

Prerequisites

Intuition

In ordinary multi-head attention (MHA) every head has its own Q, K and V. The KV cache scales with the number of heads — and becomes the memory hog at long contexts.

MQA (multi-query): all the heads share one K/V pair. The cache shrinks by the number of heads (32×), but the quality drops noticeably.

GQA (grouped-query): the heads are divided into groups that share K/V. With 32 Q heads and 8 KV groups the cache becomes 4× smaller at nearly unchanged quality.

GQA is the standard in Llama 3, Mistral and Qwen — one of the few changes that gives a large gain at almost no price.

Formal

The KV cache size:

bytes=2×L×Hkv×dhead×T×B×bytes/number\text{bytes} = 2 \times L \times H_{kv} \times d_{head} \times T \times B \times \text{bytes/number}

where LL = the layers, HkvH_{kv} = the number of KV heads (not Q heads), TT = the sequence length, BB = the batch.

An example, a 7B model (L = 32, d_head = 128, fp16, T = 8 192, B = 1):

VariantH_kvThe cache
MHA324.3 GB
GQA (4 groups)81.1 GB
MQA10.13 GB

With GQA four times as many simultaneous users fit on the same card — which translates directly into the cost per user.

What is sacrificed: expressiveness. Every group has to share the same keys and values, so heads within a group cannot specialise as freely. Empirically GQA with 4–8 groups is nearly impossible to tell from MHA, while MQA shows in the quality.

Uptraining: an MHA model can be converted to GQA by averaging the K/V projections within each group and then continuing to train on ~5 % of the original compute budget (Ainslie et al. 2023).

Code

def kv_cache_gb(layers=32, kv_heads=32, d_head=128, seq=8192, batch=1, bytes_per_number=2):
    return 2 * layers * kv_heads * d_head * seq * batch * bytes_per_number / 1e9

for name, h in (("MHA", 32), ("GQA-8", 8), ("GQA-4", 4), ("MQA", 1)):
    print(f"{name:6s} {kv_cache_gb(kv_heads=h):6.2f} GB")
# MHA     4.29 GB
# GQA-8   1.07 GB
# GQA-4   0.54 GB
# MQA     0.13 GB

# How many simultaneous users fit on an 80 GB card after the model weights (14 GB)?
for name, h in (("MHA", 32), ("GQA-8", 8)):
    print(name, int((80 - 14) / kv_cache_gb(kv_heads=h)))
# MHA 15
# GQA-8 61

Mastery means

  • Explains MHA, MQA and GQA
  • Computes the KV cache size for each variant
  • Knows what is sacrificed

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences