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

Multi-Query Attention (MQA) in C++ (LibTorch)

Multi-Query Attention (Shazeer, 2019) is a bandwidth-optimized variant of MHA. Instead of h separate K and V heads, MQA uses a single shared K and V across all Q heads. This dramatically reduces the KV cache footprint during autoregressive generation and improves throughput on memory-bandwidth-limited hardware. We'll implement it in C++ using LibTorch, with shape walkthrough.

1 · The MQA Design

Standard MHA: Q has h heads, K and V each have h heads.

MHA: Q: (batch, h, time, d_k) [h heads] K: (batch, h, time, d_k) [h heads] V: (batch, h, time, d_k) [h heads]

MQA: Q has h heads, but K and V are shared across all Q heads.

MQA: Q: (batch, h, time, d_k) [h heads] K: (batch, 1, time, d_k) [1 head, broadcast to all Q] V: (batch, 1, time, d_k) [1 head, broadcast to all Q]

When computing attention, the single K and V are implicitly replicated across the h Q heads (or equivalently, each Q head computes attention against the same K, V).

2 · KV Cache Reduction

During autoregressive generation (decoding), we cache K and V from previous tokens to avoid recomputation. With MHA, the cache grows as:

KV cache size (MHA) = 2 * (batch_size * h * max_seq_len * d_k) bytes

With MQA, the cache shrinks by a factor of h:

KV cache size (MQA) = 2 * (batch_size * 1 * max_seq_len * d_k) bytes [h times smaller]

For h=32 heads, this is a 32× reduction in memory, translating to higher throughput and longer sequence lengths on the same hardware.

3 · LibTorch Implementation

Here's the MQA module:

struct MultiQueryAttentionImpl : public torch::nn::Module { int num_q_heads, d_model, d_k; torch::nn::Linear W_q{nullptr}, W_k{nullptr}, W_v{nullptr}, W_o{nullptr}; MultiQueryAttentionImpl(int h, int d) : num_q_heads(h), d_model(d), d_k(d / h) { W_q = register_module("W_q", torch::nn::Linear(d, d)); W_k = register_module("W_k", torch::nn::Linear(d, d_k)); // single head! W_v = register_module("W_v", torch::nn::Linear(d, d_k)); // single head! 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); // Project Q to h heads: (batch, time, d_model) → (batch, time, h*d_k) torch::Tensor Q = W_q->(forward)(x); Q = Q.view({batch, time, num_q_heads, d_k}).transpose(1, 2); // Project K, V to single head: (batch, time, d_model) → (batch, time, d_k) torch::Tensor K = W_k->(forward)(x); // (batch, time, d_k) torch::Tensor V = W_v->(forward)(x); // (batch, time, d_k) // Reshape K, V to (batch, 1, time, d_k) for broadcasting K = K.unsqueeze(1); // (batch, time, d_k) → (batch, 1, time, d_k) V = V.unsqueeze(1); // Broadcast K, V: (batch, 1, time, d_k) → (batch, h, time, d_k) K = K.expand({batch, num_q_heads, time, d_k}); V = V.expand({batch, num_q_heads, time, d_k}); // Attention: Q (batch, h, time, d_k) @ K^T (batch, h, d_k, time) torch::Tensor scores = torch::matmul(Q, K.transpose(-1, -2)); scores = scores / std::sqrt(d_k); scores = torch::softmax(scores, -1); // Weighted sum: scores (batch, h, time, time) @ V (batch, h, time, d_k) torch::Tensor out = torch::matmul(scores, V); // (batch, h, time, d_k) // Transpose & reshape: (batch, h, time, d_k) → (batch, time, h*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(MultiQueryAttention);

4 · Key Shape Differences

Component MHA Shape MQA Shape Reduction
Q (batch, h, time, d_k) (batch, h, time, d_k) None
K (batch, h, time, d_k) (batch, 1, time, d_k) h× smaller
V (batch, h, time, d_k) (batch, 1, time, d_k) h× smaller
W_k d_model → d_model d_model → d_k h× fewer params
W_v d_model → d_model d_model → d_k h× fewer params

5 · Broadcasting Visualized

MQA Broadcasting: Shared K/V Across Q Heads

The single K and V head (top) is broadcast/replicated to attend with all h Q heads (bottom rows).

6 · Trade-offs

Use case: MQA shines in inference on large language models where generation is throughput-limited by memory bandwidth. It's used in LLaMA 2 70B, Falcon 40B, and other production LLMs.

7 · Grouped-Query Attention (Next Step)

Grouped-Query Attention (Part 11) is a generalization of MQA. Instead of h Q heads sharing 1 KV head, you can have G groups where each group of h/G Q heads shares one KV head. This balances cache efficiency and model capacity.