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
structGroupedQueryAttentionImpl : publictorch::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::Tensorforward(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_headstorch::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_headstorch::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^Ttorch::Tensor scores = torch::matmul(Q, K.transpose(-1, -2));
scores = scores / std::sqrt(d_k);
scores = torch::softmax(scores, -1);
// Weighted sum: scores @ Vtorch::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});
returnW_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:
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.