Grouped-query attention (GQA) keeps the full set of query heads but uses fewer key and value heads. Several query heads read the same K/V head, while each query still computes its own scores, softmax weights, and output. The architectural change narrows the K/V projections and—most importantly for autoregressive serving—stores fewer rows in the KV cache.
This places GQA on one exact axis. If the number of K/V heads equals the number of query heads, the layer is ordinary multi-head attention (MHA). If there is one K/V head, it is multi-query attention (MQA). Intermediate counts are GQA.
Follow one query from projection to decode
Change the K/V-head count to move between MHA, GQA, and MQA. Click a query lane to trace its group, then play the sequence through scoring, value mixing, compact cache storage, and the next-token append.
The head mapping
Let Hq be the number of query heads and Hkv the number of K/V heads. In the common evenly grouped layout, Hq is divisible by Hkv, and query head h reads K/V head
Each K/V head therefore serves Hq/Hkv query heads. With eight queries and two K/V heads, Q0…Q3 map to KV0, while Q4…Q7 map to KV1.
For query head h, attention remains
where M is the causal or other attention mask. Queries in one group share K and V tensors. They do not share query projections, attention-weight rows, or output vectors.
What changes in the projections
For hidden states X ∈ ℝ^(B×T×D), a GQA layer commonly uses these learned maps:
WQ ∈ ℝ^(D×Hq·d)for all query heads,WK ∈ ℝ^(D×Hkv·d)for the smaller key set,WV ∈ ℝ^(D×Hkv·d)for the smaller value set,WO ∈ ℝ^(Hq·d×D)after concatenating query-head outputs.
The reshaped tensors are
Q: [B, Hq, T, d],K: [B, Hkv, T, d],V: [B, Hkv, T, d].
Reducing Hkv narrows the K and V projection outputs, parameters, and projection work. Q and output projection widths remain tied to Hq. At one decode step, every query head still scores the available context, so the number of query-key score pairs is Hq × T; it does not fall to Hkv × T.
Why the KV cache shrinks
During autoregressive generation, each layer retains past K and V vectors. For batch size B, Nl layers, cached length T, Hkv K/V heads, head dimension d, and b bytes per element, the raw tensor payload is
The leading 2 accounts for both K and V. Relative to MHA with the same Hq, the exact payload fraction is
This is a tensor-payload calculation, not a process-memory forecast. Allocator rounding, page tables, cache blocks, quantization metadata, padding, and framework bookkeeping can change observed device memory.
Concrete model configurations
| Model | Query heads | K/V heads | Q per K/V | Relevant context mechanism |
|---|---|---|---|---|
| Llama 2 70B | 64 | 8 | 8 | 4,096-token context |
| Llama 3 8B | 32 | 8 | 4 | 8,192-token context |
| Mistral 7B v0.1 | 32 | 8 | 4 | 4,096-token sliding window |
For Llama 2 70B at 4,096 cached positions, 80 layers, head dimension 128, one sequence, and bf16, the raw GQA payload is 1.25 GiB. A hypothetical 64-K/V-head MHA cache with the same dimensions would be 10 GiB. These values exclude runtime overhead and any cache quantization.
Shape-faithful teaching implementation
The following code makes the group mapping visible. index_select creates query-aligned K/V tensors for clarity; it is not the storage strategy to copy into a production cache.
import math import torch def grouped_attention(q, k, v, mask=None): # q: [batch, Hq, query_tokens, head_dim] # k/v: [batch, Hkv, key_tokens, head_dim] hq, hkv = q.shape[1], k.shape[1] assert hq % hkv == 0 assert k.shape == v.shape group = torch.arange(hq, device=q.device) * hkv // hq k_for_q = k.index_select(1, group) # teaching expansion only v_for_q = v.index_select(1, group) scores = torch.einsum('bhtd,bhsd->bhts', q, k_for_q) scores = scores / math.sqrt(q.shape[-1]) if mask is not None: scores = scores.masked_fill(~mask, float('-inf')) weights = scores.softmax(dim=-1) return torch.einsum('bhts,bhsd->bhtd', weights, v_for_q)
A production implementation should keep the cache compact as [B, Hkv, T, d] per layer and use a kernel with native grouped-query broadcasting when available. Materializing and storing Hq repeated K/V copies forfeits the memory benefit. Rotary position embeddings, causal masking, FlashAttention-style kernels, paged caches, and cache quantization are orthogonal concerns and can be combined with GQA.
Converting an MHA checkpoint
GQA is an architectural weight-shape choice, not a safe inference-only switch. The original GQA work describes uptraining an MHA checkpoint: preserve query and output projections, mean-pool K/V heads inside each target group, and continue training. Its experiments use roughly 5% of the original pretraining compute for uptraining. That result motivates conversion, but it does not guarantee quality recovery for every model, dataset, group count, or downstream task.
Training GQA from scratch is also possible. In either route, choose Hkv using evaluation and deployment constraints rather than a universal “best” group count.
Engineering boundaries
Divisibility and partitioning
The common contiguous mapping requires Hq % Hkv == 0. Powers of two are convenient, not mathematically required. Tensor-parallel sharding may impose additional local divisibility constraints.
Cache layout
Store K/V under the K/V-head dimension, not the query-head dimension. Frameworks vary in axis order, so inspect the actual contract rather than assuming one universal [layer, batch, head, sequence, dimension] layout.
Logical broadcast versus physical repeat
A query-aligned view is useful for explaining the math. Whether expansion is zero-copy, fused, or materialized depends on the tensor strides and kernel. Profile the actual backend; do not label an expand + reshape sequence universally free.
Quality is empirical
Cache fraction and tensor shapes are exact; a generic quality label is not. Sharing can affect model quality, and the result depends on training, scale, data, evaluation, and Hkv. Report task metrics rather than an invented quality curve.
Performance is more than cache bytes
GQA can reduce K/V projection work, cache capacity, and cache bandwidth. End-to-end latency and throughput still depend on sequence lengths, batching, memory layout, kernel support, tensor parallelism, quantization, and hardware. The attention score/mix work across query heads remains.
Primary sources
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints
- Llama 2: Open Foundation and Fine-Tuned Chat Models
- Mistral 7B
- The Llama 3 Herd of Models
Related concepts
How Flash Attention, Multi-Head Attention (MHA), Grouped-Query Attention (GQA), and Multi-Query Attention (MQA) compare — algorithm vs architecture, KV-cache memory, quality trade-offs, and how to choose for production transformer inference.
Explore linear complexity attention mechanisms including Performer, Linformer, and other efficient transformers that scale to very long sequences.
Learn Multi-Query Attention (MQA), the optimization that shares keys and values across attention heads for massive memory savings.
Sliding Window Attention for long sequences: local context windows enable O(n) complexity, used in Mistral and Longformer models.
Explore sparse attention mechanisms that reduce quadratic complexity to linear or sub-quadratic, enabling efficient processing of long sequences.
Learn ALiBi, the position encoding method that adds linear biases to attention scores for exceptional length extrapolation in transformers.
