Multi-Head Latent Attention (MLA) in C++ (LibTorch)
May 2026 · 17 min read
Multi-Head Latent Attention (MLA) is a technique from DeepSeek-V2 that compresses K and V into a small latent space before caching, then reconstructs them for each attention step. This dramatically reduces KV cache size while maintaining quality through careful information bottleneck design. We'll implement the core mechanism in C++: compress X → c_KV, and reconstruct c_KV → K, V each forward pass.
1 · The Motivation
Standard attention caches K and V with shape (batch, heads, seq_len, d_k) per layer. For long contexts and many layers, this becomes a memory bottleneck. MLA takes a different approach: instead of caching K and V directly, we cache only a small bottleneck representation c_KV, and reconstruct K and V on-the-fly during decoding. The reconstruction is fast (a small projection), but the cache savings are huge.
2 · MLA Architecture
Let d_c be a small latent dimension (e.g., d_c = 128 vs d_model = 2048).
Forward pass:
c_KV = W_DKV(X) [compress to latent: (batch, time, d_c)]
K = W_UK(c_KV) [reconstruct K: (batch, time, d_model)]
V = W_UV(c_KV) [reconstruct V: (batch, time, d_model)]
Q = W_Q(X)
O = Attention(Q, K, V)
During cache (generation):
Cache c_KV instead of K, V
Each decode step: read c_KV, reconstruct K, V, compute attn
Saves 1/(h * d_k / d_c) ≈ 10-20× with clever d_c choice
3 · Optional: Low-Rank RoPE Path
DeepSeek also decomposes positional encoding into a separate low-rank path: W_KR projects X to a small d_r dimension, applies RoPE, and fuses with standard K before attention. This further decouples the frequency response of different Q heads from KV structure. For simplicity, we'll focus on the compression path but mention the pattern.
4 · LibTorch Implementation
structMultiHeadLatentAttentionImpl : publictorch::nn::Module {
int n_heads, d_model, d_k, d_c; // d_c = latent dimtorch::nn::Linear W_q{nullptr}, W_dkv{nullptr},
W_uk{nullptr}, W_uv{nullptr},
W_o{nullptr};
MultiHeadLatentAttentionImpl(int h, int d, int dc)
: n_heads(h), d_model(d), d_k(d / h), d_c(dc) {
// Q projection: full heads
W_q = register_module("W_q", torch::nn::Linear(d, d));
// Compression: X → c_KV (small latent)
W_dkv = register_module("W_dkv", torch::nn::Linear(d, d_c));
// Reconstruction: c_KV → K, V (full heads)
W_uk = register_module("W_uk", torch::nn::Linear(d_c, d));
W_uv = register_module("W_uv", torch::nn::Linear(d_c, d));
// Output projection
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);
// Step 1: Compress X to latent c_KVtorch::Tensor c_kv = W_dkv->(forward)(x); // (batch, time, d_c)// Step 2: Reconstruct K and V from latenttorch::Tensor K_full = W_uk->(forward)(c_kv); // (batch, time, d_model)torch::Tensor V_full = W_uv->(forward)(c_kv); // (batch, time, d_model)// Step 3: Standard MHA on reconstructed K, Vtorch::Tensor Q = W_q->(forward)(x);
Q = Q.view({batch, time, n_heads, d_k}).transpose(1, 2);
K_full = K_full.view({batch, time, n_heads, d_k}).transpose(1, 2);
V_full = V_full.view({batch, time, n_heads, d_k}).transpose(1, 2);
// Attentiontorch::Tensor scores = torch::matmul(Q, K_full.transpose(-1, -2));
scores = scores / std::sqrt(d_k);
scores = torch::softmax(scores, -1);
torch::Tensor out = torch::matmul(scores, V_full);
out = out.transpose(1, 2);
out = out.contiguous().view({batch, time, d_model});
returnW_o->(forward)(out);
}
// For caching during generation:torch::Tensorcompress_for_cache(torch::Tensor x) {
returnW_dkv->(forward)(x); // Cache only c_KV
}
torch::Tensorforward_cached(torch::Tensor x, torch::Tensor c_kv_cache) {
// Decode step: use cached c_KV to reconstruct K, Vtorch::Tensor Q = W_q->(forward)(x); // (batch, 1, d_model) for next tokentorch::Tensor K_full = W_uk->(forward)(c_kv_cache); // reconstruct from cachetorch::Tensor V_full = W_uv->(forward)(c_kv_cache);
// ... attend as usual ...returnW_o->(forward)(torch::zeros({1, 1, d_model})); // placeholder
}
};
TORCH_MODULE(MultiHeadLatentAttention);
5 · Shape Progression
Step
Operation
Output Shape
Memory Note
1
c_KV = W_DKV(X)
(batch, time, d_c)
Small bottleneck
2
K = W_UK(c_KV)
(batch, time, d_model)
Full K reconstructed
3
V = W_UV(c_KV)
(batch, time, d_model)
Full V reconstructed
4
Q = W_Q(X)
(batch, time, d_model)
Full Q for all heads
5
Reshape & transpose
(batch, n_heads, time, d_k)
Per-head preparation
6
scores = Q @ K^T
(batch, n_heads, time, time)
Attention weights
7
out = scores @ V
(batch, n_heads, time, d_k)
Per-head output
8
Reshape & project W_O
(batch, time, d_model)
Final output
6 · KV Cache Compression Animation
MLA Latent Compression & Reconstruction
Watch how input X is compressed to small c_KV, cached, then reconstructed on each decode step.
Trade-off: MLA trades full K, V representation for a bottleneck c_KV. If d_c is too small, information loss may hurt quality. DeepSeek tuned d_c = 128 for d_model = 2048, achieving ~97% quality vs full MHA.
8 · Extension: Dual-Stream with Low-Rank RoPE
DeepSeek-V2 adds a separate "rope stream": W_KR projects to d_r (small), applies RoPE, and concatenates or fuses with K. This allows position information to flow independently of the KV-compressed path. Implementation is similar but adds another projection pair (W_KR, W_VR).