Model Efficiency & Scaling2023intermediate10 min read
GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints
GQA: تدريب نماذج الانتباه متعدد الاستعلامات المُعمَّم من نقاط تفتيش الانتباه متعدد الرؤوس
Ainslie, J. · Lee-Thorp, J. · de Jong, M. · Zemlyanskiy, Y. · Lebrón, F. · Sanghai, S. — EMNLP
The problem
Autoregressive in Transformers is bottlenecked by : at every decoding step, all keys and values must be loaded from memory. (MQA) dramatically reduces this cost by sharing a single key-value head across all query heads, but it often degrades model quality and causes instability. Worse, many existing high-quality language models — including T5 and LLaMA — were already trained with (MHA), and retraining from scratch with MQA is prohibitively expensive.
The contribution
Two ideas. First, an recipe that converts existing multi-head checkpoints to multi-query attention using only 5% of the original compute, via mean-pooling the key and value projection matrices. Second, (GQA), which divides query heads into G groups, each sharing a single key-value head — interpolating between MHA (G=H) and MQA (G=1). Uptrained GQA-8 on T5-XXL achieves quality close to MHA with inference speed close to MQA.
The impact
GQA became the default attention mechanism in virtually every major open-weight LLM released after 2023 — Llama 2/3, Mistral, Mixtral, Gemma, and DeepSeek all use Grouped Query Attention. It proved that the quality–speed trade-off is not binary: you can keep most of multi-head quality at multi-query speed. The uptraining recipe also showed that architectural changes do not require retraining from scratch, opening a practical path for retrofitting deployed models.
Imagine a library with 64 librarians, each keeping their own personal filing cabinet of notes about every book request they have processed. When a new request arrives, all 64 cabinets must be opened and searched — that is multi-head attention. The cabinets are the , and opening them is what makes inference slow.
Multi-query attention fires all 64 librarians but keeps just one shared cabinet. Blazing fast, but a single filing system loses nuance — some requests get wrong answers.
GQA splits the 64 librarians into 8 teams of 8. Each team shares one cabinet. You still have 8 specialized cabinets (not 64, not 1), cutting storage 8× while preserving almost all the knowledge. That is the grouped-query idea: fewer notebooks, same answers.
The bottleneck: memory bandwidth in autoregressive decoding
When a generates text one token at a time, it does not recompute attention from scratch. Instead, it stores the keys and values from all previous tokens in a KV cache. At each new step, the model loads every cached key and value vector from memory to compute attention scores with the new query.
The problem is not compute — modern GPUs have plenty of FLOPs. The problem is memory bandwidth: how fast you can shuttle data between high-bandwidth memory () and the compute cores. In standard multi-head attention (MHA), the KV cache stores one key vector and one value vector per head per layer per token. For a 70-billion-parameter model with 64 heads and an 8K context, this cache alone can consume gigabytes of memory — and it must be read in full on every single decoding step.
Noam Shazeer proposed multi-query attention (MQA) in 2019: instead of H separate key-value heads, share a single key-value head across all H query heads. This cuts the KV cache by a factor of H (e.g. 64×). The speedup is dramatic, but the quality loss is real — and MQA also causes training instability, especially on long-input tasks.
Idea 1: uptraining — convert, don't retrain
Training a from scratch costs millions of dollars. If you already have a high-quality multi-head , throwing it away to retrain with MQA is wasteful. The GQA paper proposes a two-step uptraining recipe:
Step 1 — Checkpoint conversion. Take the H separate key and value projection matrices from the multi-head model and mean-pool them into the target number of heads. For MQA, that means averaging all H key matrices into one and all H value matrices into one. For GQA-G, you average each group of H/G heads. Mean-pooling preserves the most information from the original checkpoint — the paper showed it outperforms both selecting a single head and random initialization.
Step 2 — Continue pre-training. Pre-train the converted checkpoint for a small fraction α of the original training steps (the paper uses α = 0.05, i.e. 5%). This allows the model to adapt to the new weight-sharing structure. The result is a model that runs much faster at inference with only a small quality drop — and the total cost is just 5% of the original training budget.
Idea 2: Grouped Query Attention — the sweet spot
The key insight is that the jump from MHA (H key-value heads) to MQA (1 key-value head) is too aggressive. GQA introduces a tunable middle ground: divide the H query heads into G groups, and let each group share a single key-value head. This gives you G key-value heads total.
The notation GQA-G means Grouped Query Attention with G groups. Notice two special cases: GQA-1 = MQA (one group, one shared KV head) and GQA-H = MHA (each query head is its own group). Any value in between is a GQA configuration.
The KV cache size is proportional to the number of KV heads. So GQA-G reduces the cache by a factor of H/G compared to MHA. For example, if H=64 and G=8, you get an 8× reduction in KV cache — the same 8× reduction in memory bandwidth during decoding.
Why does this matter more for larger models? Because larger models have more heads, so MQA (G=1) is a more drastic cut. A model with 64 heads loses 64× capacity; one with 8 heads loses only 8×. GQA keeps the proportional reduction constant as model size grows.
The uptraining recipe in code
Simplified to show the idea — not the real implementation.
import numpy as np
def convert_mha_to_gqa(key_heads, value_heads, num_groups):
"""Convert H key/value heads into G grouped heads via mean pooling.
key_heads: (H, d_model, d_k) — one projection matrix per head
value_heads: (H, d_model, d_k)
num_groups: G — target number of KV groups
"""
H = key_heads.shape[0]
heads_per_group = H // num_groups
# Mean-pool each group of heads into one
grouped_keys = np.stack([
key_heads[g * heads_per_group : (g+1) * heads_per_group].mean(axis=0)
for g in range(num_groups)
]) # (G, d_model, d_k)
grouped_values = np.stack([
value_heads[g * heads_per_group : (g+1) * heads_per_group].mean(axis=0)
for g in range(num_groups)
]) # (G, d_model, d_k)
return grouped_keys, grouped_values
# Example: 64-head model → GQA with 8 groups
# Each group averages 8 heads into 1 → 8× KV cache reduction
# Then continue pre-training for 5% of original stepsExperiments: quality close to MHA, speed close to MQA
The authors evaluated on T5 Large and XXL with multi-head attention, plus uptrained T5-XXL with MQA and GQA-8 (α = 0.05). They tested on (CNN/DailyMail, arXiv, PubMed, MediaSum, MultiNews), translation (WMT), and (TriviaQA).
The headline result: GQA-8-XXL achieved an average score of 47.1 — nearly matching MHA-XXL's 47.2 — while running at 0.28s per sample vs. MHA-XXL's 1.51s. That is a 5.4× speedup with only 0.1 points of quality loss. Meanwhile, MQA-XXL scored 46.6 at 0.24s — faster, but with a larger 0.6-point quality drop.
Importantly, the uptrained MQA-XXL model (46.6 average, 0.24s) was both faster and higher-quality than the smaller MHA-Large model (46.0 average, 0.37s). This shows that uptraining a large model with KV sharing can beat training a smaller model with full MHA.
Ablation: what matters most?
Checkpoint conversion method: Mean-pooling outperforms selecting a single head, which in turn outperforms random initialization. The intuition is clear — averaging preserves the most information from the pretrained weights.
Uptraining proportion: GQA already achieves reasonable performance immediately after checkpoint conversion, even before any uptraining — unlike MQA, which requires uptraining to be useful at all. Both MQA and GQA benefit from 5% uptraining, with diminishing returns at 10%. This means GQA is fundamentally more robust to the conversion process.
Number of groups: Going from 1 group (MQA) to 8 groups adds very modest inference overhead, because the KV cache is already small relative to model weights. But going from 8 to 64 (full MHA) adds increasing cost. The paper chose 8 groups as a favorable middle ground, and this has become the industry default for models with 64 query heads.
Why GQA won the industry
2019
Multi-Query Attention (MQA) — Shazeer
Proposed sharing a single KV head across all query heads. Massive speedup but noticeable quality loss. Used in PaLM.
2023
GQA — Ainslie et al.
Grouped Query Attention with uptraining recipe. Quality close to MHA, speed close to MQA. Proved that the quality–speed trade-off has a sweet spot.
2023
Llama 2 adopts GQA
Meta's Llama 2 70B used GQA with 8 KV heads for 64 query heads. Set the precedent for open-weight LLMs.
2023
Mistral 7B adopts GQA
Mistral 7B combined GQA with sliding window attention. Proved GQA works at smaller scales too.
2024
DeepSeek-V3 and Llama 3 continue the trend
GQA is now the standard. DeepSeek-V3 extends the idea further with Multi-Head Latent Attention (MLA), compressing KV even more aggressively.
Broader perspective
GQA is a remarkably simple idea — just share key-value heads in groups instead of completely or not at all — but its impact is enormous precisely because it tackles the real deployment bottleneck: memory bandwidth, not FLOPs. Modern LLM serving is memory-bound during decoding, and GQA directly reduces the amount of data that must be read at each step.
The uptraining recipe is equally important. It showed the community that you can retrofit efficiency improvements onto existing models without retraining from scratch. This opened the door for a whole class of "post-hoc" optimizations — including paged attention, , and KV cache — that treat existing model weights as a starting point rather than a constraint.
GQA also interacts favorably with other efficiency techniques. Paged attention manages KV cache memory more efficiently in serving systems. reduces the FLOP overhead of attention. Quantization compresses the KV cache further. GQA is orthogonal to all of these — they stack.
CitationAinslie, Lee-Thorp, de Jong, Zemlyanskiy, Lebrón, Sanghai. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP, 2023.
Terms in this paper
- Grouped Query Attentionانتباه الاستعلام المُجمَّع
- Multi-Query Attentionانتباه الاستعلام المتعدّد
- Multi-Head Attentionالانتباه المتعدد المسارات
- KV Cacheذاكرة المفاتيح والقيم
- Inferenceالاستدلال
- Uptrainingإعادة التدريب الجزئي
- Memory Bandwidthنطاق الذاكرة
- Mean Poolingالتجميع بالمتوسط
- Checkpointنقطة التفتيش
- Attentionآلية الانتباه
- Decoderمفكّ الترميز
- Self-Attentionالانتباه الذاتي
- Cross-Attentionالانتباه التبادلي
- Fine-Tuningالضبط الدقيق