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.
MQA: Q has h heads, but K and V are shared across all Q heads.
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:
With MQA, the cache shrinks by a factor of h:
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:
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
- Pro: Dramatically smaller KV cache during generation. For long sequences, this is the dominant factor.
- Pro: Fewer parameters in W_k and W_v, slight training speedup.
- Con: All Q heads share the same K/V attention patterns. No per-head variation in "what to retrieve." In practice, the Q head projections still learn different subspaces, so this limitation is mild.
- Con: Slightly lower capacity compared to full MHA, though empirically the loss in quality is small (1–2% on some benchmarks).
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.