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

Multi-Head Latent Attention (MLA) in C++ (LibTorch)

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

struct MultiHeadLatentAttentionImpl : public torch::nn::Module { int n_heads, d_model, d_k, d_c; // d_c = latent dim torch::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::Tensor forward(torch::Tensor x) { int batch = x.size(0), time = x.size(1); // Step 1: Compress X to latent c_KV torch::Tensor c_kv = W_dkv->(forward)(x); // (batch, time, d_c) // Step 2: Reconstruct K and V from latent torch::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, V torch::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); // Attention torch::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}); return W_o->(forward)(out); } // For caching during generation: torch::Tensor compress_for_cache(torch::Tensor x) { return W_dkv->(forward)(x); // Cache only c_KV } torch::Tensor forward_cached(torch::Tensor x, torch::Tensor c_kv_cache) { // Decode step: use cached c_KV to reconstruct K, V torch::Tensor Q = W_q->(forward)(x); // (batch, 1, d_model) for next token torch::Tensor K_full = W_uk->(forward)(c_kv_cache); // reconstruct from cache torch::Tensor V_full = W_uv->(forward)(c_kv_cache); // ... attend as usual ... return W_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.

7 · Compression Ratio & Cache Savings

For d_model = 2048, d_c = 128, n_heads = 32, sequence length = 4096:

Standard MHA cache per layer: 2 × (32 heads × 4096 seq × 64 d_k × 4 bytes) = 33.6 MB MLA cache per layer: 1 × (4096 seq × 128 d_c × 4 bytes) = 2.0 MB Ratio: 33.6 / 2.0 ≈ 17× reduction Decode reconstruction cost (one step): W_UK: (128 × 2048) ≈ 0.3M params, O(128 × 2048) FLOPs [Negligible vs attention's O(seq × 2048) FLOPs]
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).