Chapter 5 β Attention Variants: GQA, MQA & MLA
βThe memory bottleneck in LLM inference isnβt compute β itβs the KV-cache. Different attention variants offer different tradeoffs between quality, memory, and throughput.β
The KV Bottleneck
In standard Multi-Head Attention (MHA), every head has its own Q, K, and V projections. During autoregressive generation, we cache K and V for every head at every layer for every past token. This KV-cache grows linearly with sequence length and is often the dominant memory cost at inference time (see Chapter 17 for the full deep dive).
For LLaMA-3-8B (MHA, 32 heads, d_k=128, 32 layers, FP16):
\[\text{KV-cache} = 2 \times 32 \times 32 \times 128 \times T \times 2\text{B} = 524{,}288 \times T \text{ bytes}\]At T = 8192 tokens: ~4 GB per request, just for KV-cache.
The key insight: not all heads need their own K and V. Query heads are diverse and important, but key and value heads can be shared without much quality loss.
Multi-Query Attention (MQA)
MQA (Shazeer, 2019) is the most aggressive sharing strategy: all query heads share a single K and V projection.
KV-cache memory reduction: HΓ smaller (32Γ for 32-head models).
Trade-off: slightly lower quality than MHA, especially on harder reasoning tasks. Used in PaLM, Falcon, StarCoder.
Grouped-Query Attention (GQA)
GQA (Ainslie et al., 2023) is the compromise: query heads are divided into G groups, and each group shares one K and V head. When G = 1, GQA = MQA. When G = H, GQA = MHA.
GQA in Practice
| Model | Query Heads (H) | KV Heads (G) | Queries per KV | KV Reduction |
|---|---|---|---|---|
| LLaMA-2-7B | 32 | 32 | 1 | 1Γ (MHA) |
| LLaMA-2-70B | 64 | 8 | 8 | 8Γ |
| LLaMA-3-8B | 32 | 8 | 4 | 4Γ |
| LLaMA-3-70B | 64 | 8 | 8 | 8Γ |
| Mistral 7B | 32 | 8 | 4 | 4Γ |
| Gemma 2 9B | 16 | 8 | 2 | 2Γ |
| Qwen 2.5 7B | 28 | 4 | 7 | 7Γ |
GQA Implementation
The key change from MHA is in the projection sizes and a repeat_interleave or expand to broadcast KV heads to match Q heads:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
class GroupedQueryAttention(nn.Module):
def __init__(self, d_model, n_heads, n_kv_heads):
super().__init__()
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.n_rep = n_heads // n_kv_heads # queries per KV head
self.d_k = d_model // n_heads
self.W_q = nn.Linear(d_model, n_heads * self.d_k, bias=False)
self.W_k = nn.Linear(d_model, n_kv_heads * self.d_k, bias=False) # smaller!
self.W_v = nn.Linear(d_model, n_kv_heads * self.d_k, bias=False) # smaller!
self.W_o = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, mask=None):
B, T, _ = x.shape
Q = self.W_q(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2) # [B, H, T, d_k]
K = self.W_k(x).view(B, T, self.n_kv_heads, self.d_k).transpose(1, 2) # [B, G, T, d_k]
V = self.W_v(x).view(B, T, self.n_kv_heads, self.d_k).transpose(1, 2) # [B, G, T, d_k]
# Expand KV heads to match Q heads: [B, G, T, d_k] β [B, H, T, d_k]
K = K.repeat_interleave(self.n_rep, dim=1)
V = V.repeat_interleave(self.n_rep, dim=1)
# Standard SDPA
attn = F.scaled_dot_product_attention(Q, K, V, is_causal=True) # [B, H, T, d_k]
out = attn.transpose(1, 2).contiguous().view(B, T, -1)
return self.W_o(out)
Parameter savings: W_k and W_v go from [d_model, d_model] to [d_model, G Γ d_k]. For LLaMA-3-8B (d=4096, H=32, G=8): each KV projection drops from 4096Γ4096 to 4096Γ1024 β 4Γ smaller.
Multi-Latent Attention (MLA)
MLA (DeepSeek-V2, 2024) takes a fundamentally different approach: instead of reducing the number of KV heads, it compresses the KV representations into a low-rank latent space.
The key advantage: at inference time, you only cache the compressed latent [B, T, d_c] instead of full K and V tensors. If d_c = 512 and d_model = 4096, this is an 8Γ compression of the KV-cache β even better than GQA in some configs.
DeepSeek-V2 also combines this with a decoupled RoPE where a separate small projection handles positional encoding, allowing the main KV compression to be position-independent.
Sliding Window Attention
Sliding Window Attention (Mistral, 2023) limits each tokenβs attention to a local window of W previous tokens instead of the full sequence. Outside the window, attention scores are masked to ββ.
- β Β· Β· Β· Β·
- β β Β· Β· Β·
- β β β Β· Β·
- β β β β Β·
- β β β β β
- β Β· Β· Β· Β·
- β β Β· Β· Β·
- β β β Β· Β·
- Β· β β β Β·
- Β· Β· β β β
KV-cache becomes fixed at W entries per head instead of growing with T. Information beyond the window can still flow through stacked layers (layer L attends to W, layer L+1βs window sees those representations, effectively reaching 2W tokens back).
Mistral 7B uses W = 4096 with alternating sliding-window and full-attention layers.
Comparison Summary
| Variant | KV per Head | KV Heads | Cache Size | Quality | Models |
|---|---|---|---|---|---|
| MHA | Full | H | 2 Γ L Γ H Γ d_k Γ T |
Best | GPT, LLaMA-1 |
| MQA | Full | 1 | 2 Γ L Γ 1 Γ d_k Γ T |
Lowest | PaLM, Falcon |
| GQA | Full | G | 2 Γ L Γ G Γ d_k Γ T |
Near-MHA | LLaMA-2/3, Mistral |
| MLA | Compressed | β | L Γ d_c Γ T |
Near-MHA | DeepSeek-V2/V3 |
| SWA | Full | H/G | 2 Γ L Γ H Γ d_k Γ W |
Good (local) | Mistral, Gemma |
Whatβs Next
Now that we understand the attention mechanisms inside transformer blocks, weβll see how these components assemble into complete decoder-only models β the dominant architecture for modern LLMs.
β Previous: Chapter 4 β SDPA & Multi-Head Attention Β· Next: Chapter 6 β Decoder-Only Models β
Last updated: April 2026