TOPIC #213Advanced 14 min read

Grouped-Query & Multi-Query Attention (GQA/MQA)

AI
AI & ML Editorial
Report an issue
Key takeawayCore Concept Summary

How shrinking the number of K/V heads — down to one (MQA) or a few groups (GQA) — slashes KV-cache size and decode memory bandwidth at near-zero quality cost, and why GQA became the default in Llama, Mistral, Gemma, and Qwen.

MHA → GQA → MQA: Shrinking the KV Side of Attention 🔻

Query heads stay plentiful; key/value heads are reduced to g groups. Fewer KV heads means a proportionally smaller KV cache and less HBM traffic per decoded token, with queries still distinct.

MHA → GQA → MQA: Shrinking the KV Side of Attention 🔻
100%
Touchpad: Pinch to zoom • Drag to pan
Rendering visual architecture flowchart...

01.The Problem: Generation Is Slow Because of the Filing Cabinet, Not the Math

Forget training for a moment — this bottleneck lives entirely in serving.

When a model writes text one token at a time (autoregressive generation), each new token must compare itself — via attention — against every previous token. The K and V vectors of all those previous tokens are stored so they never have to be recomputed. That store is the KV cache.

Per generated token, the GPU must read the whole cache: every layer, every KV head, every position. The arithmetic for one token is trivial (a single query row!), but the reading streams gigabytes.

Insight

So what actually limits how fast and how many users a model can serve?

Not FLOPs. Memory bandwidth. Generation is memory-bandwidth-bound, and the KV cache is the cargo being hauled.

How big is the cargo? The formula is five factors multiplied:

KV bytes per token = 2 × layers × kv_heads × head_dim × bytes_per_element × seq_len

Llama-2 70B (80 layers, head_dim 128, FP16), full multi-head attention with 64 KV heads:

  • 2 × 80 × 64 × 128 × 2 bytes ≈ 2.5 MB per token
  • one 128K-token context → 2.5 MB × 128,000 ≈ ~320 GB of cache. Absurd — no single GPU comes close.

Now the same model as it actually ships, with GQA-8 (only 8 KV heads):

  • ≈ 0.31 MB per token → ~40 GB for one 128K sequence. Still enormous for batching — which is exactly why PagedAttention, cache quantization, and sliding-window layers stack on top.

And the observation that defines this whole topic: the cache size scales with the number of KV heads — a number you are free to shrink.

02.The Idea in Plain Words: Many Questions, Fewer Reference Shelves

Attention has three roles: each query head asks "what should I look at?", and key/value heads supply the things to look at. Multi-head attention (MHA) gives every query head its own private K and V head — h of each.

GQA/MQA in one line

Insight

Keep all the query heads — they are cheap, one row each — but make groups of them SHARE one K/V head, because the KV cache is what you pay to store and re-read forever.

The spectrum, parameterized by g = number of K/V heads:

  • MHA: g = h — every query has its own KV head. Maximum expressiveness, maximum cache.
  • GQA (Grouped-Query): g in between — partition the h query heads into g groups; each group attends to its own shared KV head. E.g. 32 queries / 8 KV = 4:1.
  • MQA (Multi-Query): g = 1 — all query heads share a single K/V head. Cache is h× smaller. Note MQA is just GQA's extreme case, not a rival design.

Queries stay distinct (each still computes its own attention pattern with its own weights); only the supply side is compressed. Fewer KV heads also shrink the K/V projection weight matrices themselves — a minor bonus.

03.A Worked Example: The 4× Shrink, Line by Line

Llama-3-8B style: 32 query heads, 32 layers, head_dim 128, FP16 (2 bytes).

MHA (32 KV heads):

2 × 32 layers × 32 kv_heads × 128 × 2 bytes = 524,288 bytes ≈ 512 KB per token

At 8K context: 512 KB × 8,192 ≈ 4 GB of KV cache — per request.

GQA (8 KV heads, groups of 4):

2 × 32 × 8 × 128 × 2 = 131,072 bytes ≈ 128 KB per token

At 8K context: 128 KB × 8,192 ≈ 1 GB. Same 4× ratio, one-quarter the traffic, identical query heads.

MQA (1 KV head): ≈ 16 KB per token, 400 MB at 8K.

Why decode cares so much: the batch-size-1 decode step computes 8 heads-worth of dot products (tiny) but must read the entire cache (all of it). At ~3 TB/s on an H100, shaving 4 GB to 1 GB per request directly multiplies how many concurrent long-context requests fit and stream. Cache bytes are the rent; fewer KV heads lowers it proportionally.

04.Visual Intuition: Funnel the Supply Side

code
  MHA (8 queries, 8 KV heads)          GQA (8 queries, 4 KV heads)
  Q1 Q2 Q3 Q4 Q5 Q6 Q7 Q8              Q1 Q2 Q3 Q4 Q5 Q6 Q7 Q8
  │  │  │  │  │  │  │  │                │ │  │ │  │ │  │ │
  K1 V1 K2 V2 K3 V3 ...one each         K1V1 K2V2 K3V3 K4V4
  cache slots per token: 8              each serves a GROUP of 2
                                        cache slots per token: 4

  MQA (8 queries, 1 KV head)
  Q1..Q8 all read the SAME K1 V1
  cache slots per token: 1

The queries fan IN (their number never drops — expressiveness stays); the keys/values fan OUT into fewer slots — and it is the KV side, not the Q side, that gets cached and re-streamed on every generated token. That is why shrinking only it is nearly free.

05.The Analogy: The Newsroom With Shared Reference Desks

Carry one picture: a newsroom of 32 journalists (queries) writing stories from the reference archive (keys/values).

  • MHA: each journalist has a private archive wing. Perfect for depth, but the building is 32 copies of the reference material, and every time someone re-checks the files they haul an entire wing down the hall (bandwidth).
  • MQA: one shared archive for the whole floor. Storage collapses, hauls are tiny — but now 32 journalists crowd the same shelves: the reference material gets generic, blurry about each beat, and quality suffers. MQA (Shazeer, 2019, arXiv 1911.02150) made exactly this extreme bet — all query heads share a single K/V head — originally to make giant autoregressive decoding cheap in TPU serving. It proved workable, but the reported perplexity degradation made many practitioners judge it too risky, and dense multi-head attention stayed the default for years.
  • GQA: archive wings per department — foreign desk, business desk, science desk each keep one reference shelf for their 4 reporters. Almost all the storage savings of MQA while each department still has a specialized, less-crowded supply.

The empirical surprise of the GQA paper: going from 32 private wings to 8 departmental shelves cost essentially nothing in quality — the "quality curve is flat until g gets very small." And the retrofit trick ("up-training") is like re-tagging the existing 32 wings' holdings into 8 shelves over a weekend plus a short re-training period, rather than rebuilding the newsroom (full retraining).

06.GQA: The Sweet Spot, and Retrofitting From MHA

Grouped-Query Attention (Ainslie et al., Google, 2023, arXiv 2305.13245) interpolates between the two classics: partition the h query heads into g groups, each group attending to its own shared K/V head. g = h recovers MHA; g = 1 recovers MQA. Two findings made it the industry default:

  1. The quality curve is flat until g gets very small. 8 groups on 64 heads costs essentially nothing versus MHA and is markedly more robust than MQA — the "upsample then train" experiments showed GQA matching MHA quality while inheriting MQA-class cache savings.
  2. You can retrofit. Their "up-training" recipe takes an existing MHA checkpoint, averages K/V head weights within the planned groups (or up-projects the averaged weights back), then fine-tunes briefly. Llama 2 70B and Llama 3 shipped exactly this way — converted from a 34B-style dense model. No from-scratch retrain.

The mechanics at run time are cheap: query projections and the attention math keep the full head count; only the K/V projections output g × head_dim, broadcast to the queries within each group. Kernel support is universal — FlashAttention, cuDNN, and FlashInfer all handle g ≠ h natively (they just repeat each KV head across its group inside the kernel).

07.What the Big Models Actually Chose (2024-2026)

By 2025, GQA is the de-facto standard configuration in open-weight transformers, with group counts tuned to context ambition:

  • Mistral 7B (2023): MQA — the proof that a single-KV-head design could win production at scale (an early vLLM sweet spot).
  • Llama 2 70B / Llama 3 8B & 70B: GQA with 8 KV heads (against 32 and 64 query heads respectively).
  • Gemma 2 / Gemma 3, Qwen2.5, OLMo 2: GQA (typical query-to-KV ratios 4:1 to 8:1).
  • Mixtral 8x7B: inherits Mistral's MQA; DeepSeek-V3: diverges with MLA instead.

The practical deployment impact: a served 70B-class model with GQA-8 sustains 5-10× the concurrent long-context batch of its MHA twin on identical HBM — or equivalently drops p99 decode latency because cache reads shrink. Combined with KV-cache quantization and PagedAttention (both elsewhere in this phase), GQA is one leg of the stool that made 128K-context APIs economically viable.

python— GQA projection shapes for 32 query heads but 8 KV heads (Llama-3-8B style)
import torch.nn as nn

n_q_heads, n_kv_heads, head_dim = 32, 8, 128
d_model = n_q_heads * head_dim           # 4096

Wq = nn.Linear(d_model, n_q_heads * head_dim, bias=False)   # 4096 -> 4096
Wk = nn.Linear(d_model, n_kv_heads * head_dim, bias=False)  # 4096 -> 1024  <-- 4x smaller
Wv = nn.Linear(d_model, n_kv_heads * head_dim, bias=False)  # 4096 -> 1024  <-- 4x smaller

# At attention time each KV head is broadcast (repeat_interleave) to its
# group of 4 query heads; FlashAttention does this implicitly in-kernel.
# KV cache per token: 2 * n_layers * 8 * 128 * bytes  (4x smaller than MHA)

Architectural Trade-offs & Production Realities

Architectural Advantages

  • KV cache and decode HBM traffic shrink by h/g (e.g., 8× for GQA-8 on 64 heads) — direct throughput and cost win.
  • Near-identical quality to MHA; far more robust than MQA across tasks.
  • Retrofit path from existing MHA checkpoints (up-training) avoids full retraining.

Trade-offs & Constraints

  • Slight expressiveness loss vs MHA on some tasks (reasoning-heavy evals can show a hair of regression at aggressive g).
  • K/V projection capacity shrinks too — very low g can underfit on huge vocab/model scales.
  • One more config dimension (n_kv_heads) to tune; naive implementations need group-broadcast logic.
Production Implementation in Big Tech
Meta Llama 3 / Mistral• KV-cache shrink enabling 128K-context serving

Llama 3 (8B and 70B) uses GQA with 8 KV heads against 32/64 query heads, cutting cache traffic 4x-to-8x versus MHA designs of the same era; Mistral 7B chose full MQA, letting early vLLM deployments serve long contexts on a single 24-32GB GPU. Both choices underpin the 2024-2026 open-model serving stack.

Staff+ Engineering Takeaways

  • Decode is memory-bandwidth-bound on the KV cache; KV bytes scale with the number of KV heads, not query heads.
  • MQA shares a single K/V head across all queries (1911.02150); GQA groups queries under g shared KV heads (2305.13245).
  • GQA-8 on 64 query heads keeps ~MHA quality with 8× smaller cache and dramatically higher concurrent batch capacity.
  • Existing MHA checkpoints can be up-trained to GQA — the route Llama 2 70B/Llama 3 took.
  • Default in 2024-2026 open models: GQA everywhere (Mistral earlier used MQA); DeepSeek diverges with MLA latent compression.

Topic Knowledge Check

Exercise 1 of 3 • Test your architectural comprehension.

Exercise 1 of 30 answered
1

A model has 64 query heads and 8 KV heads (GQA with 8 groups). Compared to full MHA with 64 KV heads, how large is its KV cache per token?

Rate This Architecture ChapterFeedback & Rating

How clear and actionable was this distributed systems breakdown?