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

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:

MHA(X) = Concat(head_1, …, head_h) · W_O where head_i = Attention(X · W_Q^{(i)}, X · W_K^{(i)}, X · W_V^{(i)})

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:

X · W_Q: (batch, time, d_model) ↓ reshape to: (batch, time, h, d_k) [d_model = h * d_k] ↓ transpose to: (batch, h, time, d_k) ↓ attention ↓ transpose back: (batch, time, h, d_k) ↓ reshape: (batch, time, d_model) ↓ project O: (batch, time, d_model)

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:

struct MultiHeadAttentionImpl : public torch::nn::Module { int num_heads, d_model, d_k; torch::nn::Linear W_q{nullptr}, W_k{nullptr}, W_v{nullptr}, W_o{nullptr}; MultiHeadAttentionImpl(int h, int d) : num_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)); W_v = register_module("W_v", torch::nn::Linear(d, d)); 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 to Q, K, V: each (batch, time, d_model) torch::Tensor Q = W_q->(forward)(x); torch::Tensor K = W_k->(forward)(x); torch::Tensor V = W_v->(forward)(x); // Reshape: (batch, time, d_model) → (batch, time, h, d_k) Q = Q.view({batch, time, num_heads, d_k}); K = K.view({batch, time, num_heads, d_k}); V = V.view({batch, time, num_heads, d_k}); // Transpose: (batch, time, h, d_k) → (batch, h, time, d_k) Q = Q.transpose(1, 2); K = K.transpose(1, 2); V = V.transpose(1, 2); // Attention per head: (batch, h, time, d_k) × (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 values: (batch, h, time, time) × (batch, h, time, d_k) torch::Tensor out = torch::matmul(scores, V); // (batch, h, time, d_k) // Transpose back: (batch, h, time, d_k) → (batch, time, h, d_k) out = out.transpose(1, 2); // Reshape merge: (batch, time, h, d_k) → (batch, time, d_model) out = out.contiguous().view({batch, time, d_model}); // Final projection return W_o->(forward)(out); } }; TORCH_MODULE(MultiHeadAttention);

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:

Scaling: With h = 8 heads and d_model = 512, each head sees d_k = 64 dimensions. After attention in all heads and concatenation, we project back through W_O to fuse head outputs before passing to the next layer.

7 · Causal Masking & Dropout

To add causal attention (decoder), insert a mask before softmax:

torch::Tensor causal_mask(int time) { return torch::triu(torch::ones({time, time}), 1) .to(torch::kFloat32) * -1e9; } // In forward, after scores scaling: torch::Tensor mask = causal_mask(time).unsqueeze(0).unsqueeze(0); scores = scores + mask;