← All Posts
From Scratch · C++ · Transformers · Attention Variants· Part 11 of 13

Grouped-Query Attention (GQA) in C++ (LibTorch)

Grouped-Query Attention (Ainslie et al., 2023) generalizes Multi-Query Attention by introducing a configurable number of KV heads (n_kv_heads ≤ n_q_heads). Multiple Q head groups share a single KV head each. This balances the model capacity of full MHA against the cache efficiency of MQA. We'll implement it in C++ with torch::repeat_interleave for clean grouping.

1 · The Spectrum: MHA → GQA → MQA

Think of attention mechanisms as a spectrum:

MHA (h heads each): n_q_heads = h, n_kv_heads = h → each Q head gets its own K, V GQA (groups): n_q_heads = h, n_kv_heads = g (g | h, g < h) → h/g Q heads share each KV head MQA (shared): n_q_heads = h, n_kv_heads = 1 → all Q heads share 1 K, V

With GQA, the KV cache size scales with n_kv_heads, not n_q_heads. For n_q_heads=32, n_kv_heads=8, we save a 4× factor compared to MHA.

2 · Shape and Replication Strategy

After projection:

Q: (batch, n_q_heads, time, d_k) K: (batch, n_kv_heads, time, d_k) V: (batch, n_kv_heads, time, d_k)

To compute attention with shape alignment, we replicate K and V using repeat_interleave:

groups_per_kv = n_q_heads / n_kv_heads K_expanded: (batch, n_q_heads, time, d_k) [each KV head replicated groups_per_kv times] V_expanded: (batch, n_q_heads, time, d_k)

Now all shapes align: Q @ K_expanded^T and scores @ V_expanded work without broadcasting tricks.

3 · LibTorch Implementation

struct GroupedQueryAttentionImpl : public torch::nn::Module { int n_q_heads, n_kv_heads, d_model, d_k; torch::nn::Linear W_q{nullptr}, W_k{nullptr}, W_v{nullptr}, W_o{nullptr}; GroupedQueryAttentionImpl(int n_q, int n_kv, int d) : n_q_heads(n_q), n_kv_heads(n_kv), d_model(d), d_k(d / n_q) { TORCH_CHECK(n_q % n_kv == 0, "n_q_heads must be divisible by n_kv_heads"); W_q = register_module("W_q", torch::nn::Linear(d, d)); W_k = register_module("W_k", torch::nn::Linear(d, n_kv_heads * d_k)); W_v = register_module("W_v", torch::nn::Linear(d, n_kv_heads * d_k)); W_o = register_module("W_o", torch::nn::Linear(d, d)); } torch::Tensor forward(torch::Tensor x) { int batch = x.size(0), time = x.size(1); int group_factor = n_q_heads / n_kv_heads; // Project Q to n_q_heads torch::Tensor Q = W_q->(forward)(x); Q = Q.view({batch, time, n_q_heads, d_k}).transpose(1, 2); // Project K, V to n_kv_heads torch::Tensor K = W_k->(forward)(x); // (batch, time, n_kv_heads*d_k) K = K.view({batch, time, n_kv_heads, d_k}).transpose(1, 2); torch::Tensor V = W_v->(forward)(x); V = V.view({batch, time, n_kv_heads, d_k}).transpose(1, 2); // Replicate K, V: (batch, n_kv_heads, time, d_k) → (batch, n_q_heads, time, d_k) K = torch::repeat_interleave(K, group_factor, 1); V = torch::repeat_interleave(V, group_factor, 1); // Attention: Q (batch, n_q_heads, time, d_k) @ K^T torch::Tensor scores = torch::matmul(Q, K.transpose(-1, -2)); scores = scores / std::sqrt(d_k); scores = torch::softmax(scores, -1); // Weighted sum: scores @ V torch::Tensor out = torch::matmul(scores, V); // (batch, n_q_heads, time, d_k) // Reshape merge: (batch, n_q_heads, time, d_k) → (batch, time, d_model) out = out.transpose(1, 2); out = out.contiguous().view({batch, time, d_model}); return W_o->(forward)(out); } }; TORCH_MODULE(GroupedQueryAttention);

4 · Shape Walkthrough

Operation Input Output Notes
Q = W_q(x) (batch, time, d_model) (batch, n_q_heads, time, d_k) Full n_q_heads
K = W_k(x) (batch, time, d_model) (batch, n_kv_heads, time, d_k) Only n_kv_heads (n_q_heads / group_factor)
repeat_interleave(K, group_factor, 1) (batch, n_kv_heads, time, d_k) (batch, n_q_heads, time, d_k) Replicate along head dim
scores = Q @ K^T (batch, n_q_heads, time, d_k) × (batch, n_q_heads, d_k, time) (batch, n_q_heads, time, time) Aligned shapes
out = scores @ V (batch, n_q_heads, time, time) × (batch, n_q_heads, time, d_k) (batch, n_q_heads, time, d_k) Weighted sum per head
reshape to d_model (batch, n_q_heads, time, d_k) (batch, time, d_model) Merge heads & project

5 · Grouping Visualization

GQA Groups: Q Heads Sharing KV Heads

Watch how each KV head (top row) is replicated to multiple Q head groups (bottom grid). Colors show group membership.

6 · Parameter & Memory Savings

For a layer with n_q_heads = 32, n_kv_heads = 8, d_k = 64, d_model = 2048:

MHA params (W_q + W_k + W_v): (2048 × 2048) + (2048 × 2048) + (2048 × 2048) = 12.6M GQA params: (2048 × 2048) + (2048 × 512) + (2048 × 512) = 6.3M [50% reduction] MQA params: (2048 × 2048) + (2048 × 64) + (2048 × 64) = 4.2M [67% reduction] KV cache (during generation, batch=1, seq=4096): MHA: 2 × 32 × 4096 × 64 = 16.8M scalars GQA: 2 × 8 × 4096 × 64 = 4.2M scalars [75% reduction vs MHA] MQA: 2 × 1 × 4096 × 64 = 0.5M scalars [97% reduction vs MHA]

7 · Configuration Guide

Choose n_kv_heads based on your constraints:

Empirical note: On many LLMs, GQA with n_kv_heads = n_q_heads / 4 achieves >95% of full MHA quality while cutting KV cache by 75%. This is often the Pareto-optimal choice.