Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

Explain Grouped-Query Attention (GQA) vs Multi-Head Attention—memory and quality differences?


Eight query heads reading two cached K/V headsq0q1q2q3q4q5q6q701234567group 1reads KV head 0group 2reads KV head 1Only 2 K/V heads are stored: a 4x smaller cache, the same number of query heads.
GQA keeps every query head asking its own question and shares only what gets cached, so memory falls while attention FLOPs stay put.

What you need to know

Three layouts of the same idea

Text
MHA:  n_q = 32, n_kv = 32   → cache ∝ 32 headsGQA:  n_q = 32, n_kv = 8    → cache ∝ 8 heads   (4× smaller)MQA:  n_q = 32, n_kv = 1    → cache ∝ 1 head    (32× smaller)

MQA came from Shazeer (2019), "Fast Transformer Decoding: One Write-Head is All You Need". GQA came from Ainslie et al. (2023), who also showed you can convert an existing MHA checkpoint — mean-pool the K/V heads within each group, then "uptrain" with about 5% of the original pretraining compute — and recover quality close to MHA while keeping most of MQA's speed.

What changes and what does not

  • KV cache — shrinks by n_q / n_kv. This is the main win.
  • Decode speed — decode is memory-bound, and at long context the cache can be a large part of what each step reads, so a smaller cache means faster steps and bigger batches.
  • Attention FLOPs — unchanged. Each query head still computes a full row of scores against keys.
  • Parameters — W_k and W_v shrink. In Llama-3-8B each is 4096 × 1024 instead of 4096 × 4096, saving about 25M parameters per layer: small next to the cache saving.

How it looks in code

Python
import torchB, T, n_q, n_kv, hd = 1, 6, 32, 8, 128q = torch.randn(B, n_q, T, hd)k = torch.randn(B, n_kv, T, hd)          # only 8 K heads are cachedv = torch.randn(B, n_kv, T, hd)group = n_q // n_kv                      # 4 query heads per K/V headk_exp = k.repeat_interleave(group, dim=1)   # (1, 32, 6, 128) for the mathv_exp = v.repeat_interleave(group, dim=1)out = torch.nn.functional.scaled_dot_product_attention(q, k_exp, v_exp, is_causal=True)print(out.shape, k.numel() / k_exp.numel())   # torch.Size([1, 32, 6, 128]) 0.25

Only the 8-head k and v are stored. Query heads 0–3 read KV head 0, heads 4–7 read KV head 1, and so on. Recent PyTorch versions accept enable_gqa=True in scaled_dot_product_attention, so optimised kernels do this sharing without materialising the expanded copy.

Where the field went next

GQA with 8 KV heads became the default for open models from Llama 2 70B onward. DeepSeek-V2 (2024) introduced multi-head latent attention (MLA), which caches one small compressed vector per token and reconstructs K and V from it; the paper reports a 93.3% smaller KV cache than their earlier 67B dense model. Some models also share one KV cache across neighbouring layers.

A real-life example

A team serves a chatbot on a 70B model across 8 H100s (640 GB). The fp16 weights take about 140 GB, leaving roughly 500 GB for cache. Llama-3-70B has 80 layers and head_dim 128.

Text
MHA (64 KV heads): 2 × 80 × 64 × 128 × 2 bytes = 2.5 MiB per token → 20 GiB per 8k chatGQA (8 KV heads):                               320 KiB per token → 2.5 GiB per 8k chat

With MHA the node holds about 25 full-length conversations; with GQA, about 200. For a chatbot with a million daily users, that 8× difference is the difference between a cluster of nodes and a building's worth of them.

Follow-up questions to expect

  • "Does GQA reduce attention FLOPs?" — Hardly. Every query head still scores every key; only the K/V projections and the cache get smaller.
  • "How do you choose the number of KV groups?" — It is a quality-versus-memory knob; 8 is common because it also divides evenly across 8-way tensor parallelism.
  • "Why is MQA worse?" — All heads must read the same keys and values, which limits how differently heads can attend, and the original GQA paper also reports MQA being less stable to train.