← Back to Table of Contents

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.

MQA β€” Single KV Head Shared Across All Query Heads
Q head 1
Q head 2
Q head 3
…
Q head H
↓ all attend using ↓
K head 1 (shared)   V head 1 (shared)
MQA Tensor Shapes
Q
[ B, H, T, d_k ]
H independent query heads
K, V
[ B, 1, T, d_k ]
1 shared KV head β†’ broadcast to H

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 β€” Groups of Query Heads Share KV Heads
Group 1
Q heads 1–4 share KV head 1
Group 2
Q heads 5–8 share KV head 2
Group 3
Q heads 9–12 share KV head 3
Group 4
Q heads 13–16 share KV head 4
Example: H=16 query heads, G=4 KV heads β†’ 4 queries per KV head
GQA Tensor Shapes
Q
[ B, H, T, d_k ]
H = 32 query heads
K, V
[ B, G, T, d_k ]
G = 8 KV heads β†’ each shared by H/G = 4 queries

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.

MLA β€” Low-Rank Compression of KV
Input x: [B, T, d_model]
Down-project: W_dkv [d_model, d_c] β†’ compressed KV: [B, T, d_c]
Cache only compressed: [B, T, d_c] d_c β‰ͺ d_model (e.g., 512 vs 4096)
Up-project at decode time: [B, T, d_c] β†’ K, V: [B, H, T, d_k]
Standard SDPA with reconstructed K, V

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 βˆ’βˆž.

Sliding Window vs Full Causal Attention (W=3)
Full Causal (T=5)
  • βœ“ Β· Β· Β· Β·
  • βœ“ βœ“ Β· Β· Β·
  • βœ“ βœ“ βœ“ Β· Β·
  • βœ“ βœ“ βœ“ βœ“ Β·
  • βœ“ βœ“ βœ“ βœ“ βœ“
Sliding Window (W=3)
  • βœ“ Β· Β· Β· Β·
  • βœ“ βœ“ Β· Β· Β·
  • βœ“ βœ“ βœ“ Β· Β·
  • Β· βœ“ βœ“ βœ“ Β·
  • Β· Β· βœ“ βœ“ βœ“

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

Side-by-side comparison of MHA, GQA, and MQA showing Q, K, V head counts and KV cache sizes
Multi-Head Attention vs Grouped-Query Attention vs Multi-Query Attention β€” head configurations and cache implications
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
KV-Cache Memory vs Sequence Length
MHA (H=32)
Linear growth: 512 KB/token for LLaMA-3-8B scale. 4 GB at 8K tokens.
GQA (G=8)
4Γ— smaller: 128 KB/token. 1 GB at 8K tokens. Near-MHA quality.
MQA (G=1)
32Γ— smaller: 16 KB/token. 128 MB at 8K. Some quality tradeoff.
MLA
Compressed latent: much smaller than GQA with comparable quality. DeepSeek-specific.

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