Multi-Head Attention (MHA) in C++ (LibTorch)
Multi-Head Attention (MHA) is the workhorse of modern transformers. This post focuses on the single-projection-per-type design: one $W_Q$, $W_K$, $W_V$ matrix of shape (d_model, d_model) each, rather than separate matrices for each head. We reshape and transpose to split the output into h parallel heads, compute attention in each, and merge back. This is the standard and most efficient formulation.
1 · The MHA Formula
Given input X of shape (batch, time, d_model), MHA produces output of the same shape:
The key insight is that instead of separate weight matrices for each head, we use one flat W_Q of size (d_model, d_model), compute X · W_Q (shape: (batch, time, d_model)), and then reshape and transpose to separate the heads. This is mathematically equivalent but requires fewer distinct parameter tensors and enables better memory layout for parallelism.
2 · Shape Transformations
The reshape-transpose pipeline is the core of MHA:
Each head operates in parallel on its own (batch, time, d_k) slice. After attention in all heads, we concatenate (reshape merge) and project through W_O.
3 · LibTorch Implementation
Here is a complete MultiHeadAttention module:
4 · Shape Walkthrough Table
| Step | Operation | Input Shape | Output Shape |
|---|---|---|---|
| 1 | Q = W_q(x) |
(batch, time, d_model) | (batch, time, d_model) |
| 2 | Q.view(batch, time, h, d_k) |
(batch, time, d_model) | (batch, time, h, d_k) |
| 3 | Q.transpose(1, 2) |
(batch, time, h, d_k) | (batch, h, time, d_k) |
| 4 | Q @ K.T |
(batch, h, time, d_k) × (batch, h, d_k, time) | (batch, h, time, time) |
| 5 | softmax(scores) |
(batch, h, time, time) | (batch, h, time, time) |
| 6 | scores @ V |
(batch, h, time, time) × (batch, h, time, d_k) | (batch, h, time, d_k) |
| 7 | transpose(1, 2) |
(batch, h, time, d_k) | (batch, time, h, d_k) |
| 8 | view(batch, time, d_model) |
(batch, time, h, d_k) | (batch, time, d_model) |
| 9 | W_o(out) |
(batch, time, d_model) | (batch, time, d_model) |
5 · Shape Pipeline Animation
MHA Reshape & Transpose Flow
Watch the shape pipeline: input → project → reshape to heads → transpose → attention → transpose → merge → output.
6 · Why Single-Flat Projections?
An earlier design (sometimes called "stacked heads") creates separate W_Q matrices for each head and concatenates. This requires h times more distinct parameters and more complex registration logic. The flat-project-then-reshape approach:
- Requires only one W_Q, W_K, W_V, W_O per layer (not per head).
- Naturally maps to single GEMM operations, enabling GPU acceleration.
- Aligns with modern framework optimizations (cuBLAS, oneDNN) that handle large, batched matrix multiplies efficiently.
7 · Causal Masking & Dropout
To add causal attention (decoder), insert a mask before softmax: