MultiHeadAttention in C++ (LibTorch)
For the theory behind multi-head attention, including why we need multiple heads, the split-attend-concat pipeline, and what different heads learn, see the dedicated Multi-Head Attention blog. This page walks through the full C++ implementation using LibTorch, line by line.
1 · What Does MultiHeadAttention Do?
At a high level, multi-head attention does three things:
- Runs multiple SingleHeadAttention modules in parallel. Each head receives the same input tensor but processes it through its own learned Q, K, V projections.
- Concatenates all head outputs along the last dimension. If you have 8 heads, each producing a vector of size 32, the concatenated result is a vector of size 256.
- Projects the concatenated output through a linear layer. This learned projection mixes information across heads and maps the result back to the model's embedding dimension.
Our implementation follows this exactly. The SHA_group vector holds the heads, and the projection_layer is WO.
Input and output shapes
The module takes in a 3D tensor and produces a 3D tensor of the same shape:
This shape preservation is essential. In a transformer, the output of multi-head attention feeds into a residual connection: x = x + MultiHeadAttention(x). If the shapes did not match, the addition would fail.
2 · The Class Structure
Here is the complete class definition, exactly as written:
Let us break down the key design decisions.
Member variables
Three integers (num_heads, head_size, n_embd) store the configuration. The vector SHA_group holds shared_ptrs to each SingleHeadAttention head — shared pointers are required because LibTorch's register_module needs shared ownership. The projection_layer is initialized to nullptr and constructed in the constructor body.
The constructor: registering heads in a loop
The constructor uses a loop to create and register each head:
Three things happen on each iteration:
- Create the head:
make_shared<SingleHeadAttention>(n_embd, head_size)allocates a new SingleHeadAttention on the heap. Each head gets its own Q, K, V weight matrices of shape(n_embd, head_size). - Register the head:
register_module("SHA_" + to_string(head_idx), curr_head)adds this head to the LibTorch module tree. The name is unique per head:"SHA_0","SHA_1","SHA_2", etc. This registration is what makes the head's parameters visible to the optimizer. Without it, the Q, K, V weights inside the head would never be updated during training. - Store the head:
SHA_group.push_back(curr_head)keeps a reference in the vector for use during the forward pass.
torch::nn::Module with its own parameters. If you skip registration, those parameters become invisible to the optimizer and to model->to(torch::kCUDA).
to_string: std::to_string lives in <string>, but you don't need an extra include — LibTorch's umbrella header <torch/torch.h> transitively pulls in <string> (and most of the standard library headers you'll need). If you ever use a minimal include like <torch/nn/module.h> instead, add #include <string> explicitly.
The projection layer
The projection layer maps from head_size * num_heads to n_embd. Since these are equal by design, it is an n_embd → n_embd linear transformation. Unlike the Q, K, V layers which omit bias, the projection layer uses the default which includes bias.
3 · The Key Relationship: head_size * num_heads = n_embd
This is the central constraint that makes multi-head attention work. The embedding dimension n_embd is divided equally among the heads. Each head operates in a subspace of dimension head_size, and when their outputs are concatenated, we recover the full n_embd dimension.
Some more examples to build intuition:
| n_embd | num_heads | head_size | Each head output | Concatenated |
|---|---|---|---|---|
256 | 4 | 64 | (B, S, 64) | (B, S, 256) |
256 | 8 | 32 | (B, S, 32) | (B, S, 256) |
512 | 8 | 64 | (B, S, 64) | (B, S, 512) |
768 | 12 | 64 | (B, S, 64) | (B, S, 768) |
1024 | 16 | 64 | (B, S, 64) | (B, S, 1024) |
n_embd = 256 and num_heads = 7, then head_size = 256 / 7 = 36.57, which is not an integer. Always choose num_heads that evenly divides n_embd.
n_embd-dimensional input and projects it down to head_size dimensions using its own learned weight matrices. The splitting happens in the output space, not the input space.
4 · The Forward Pass Line by Line
Here is the complete forward method:
Let us walk through every line.
Line 1: The head loop
Each iteration retrieves the head at head_idx, calls forward(x), and pushes the result — shape (B, S, head_size) — into the SHA_outputs vector. Every head receives the exact same x; the different Q, K, V weights are what produce different outputs.
Line 2: Concatenation
torch::cat joins tensors along a specified dimension. -1 means the last dimension (the feature dimension). For 8 heads each producing (B, S, 32), the result is (B, S, 256).
Line 3: Projection
The projection layer applies a (n_embd, n_embd) weight matrix plus bias. This is the step where information from different heads gets mixed — each output value is a weighted sum across all concatenated head dimensions.
Line 4: Return
The projected output, with shape (B, S, n_embd), is returned. This feeds into the next layer of the transformer: typically a residual addition followed by layer normalization and then a feed-forward network.
5 · Interactive Animation
Step through the forward pass with 3 heads. Watch how each head produces its own output, how the outputs are concatenated, and how the projection layer maps back to n_embd.
Multi-Head Attention: 3 Heads
n_embd = 96, num_heads = 3, head_size = 32, batch = 1, seq_len = 4
7 · Parameter Count
Using concrete numbers: n_embd = 256, num_heads = 8, head_size = 32.
The general formula simplifies to 4 * n_embd² + n_embd, which is independent of the number of heads.
| Model | n_embd | num_heads | head_size | MHA Params |
|---|---|---|---|---|
| GPT-2 Small | 768 | 12 | 64 | 2,360,064 |
| GPT-2 Medium | 1024 | 16 | 64 | 4,195,328 |
| GPT-2 Large | 1280 | 20 | 64 | 6,554,880 |
| GPT-2 XL | 1600 | 25 | 64 | 10,241,600 |
8 · What Comes Next
We now have the two core attention building blocks: SingleHeadAttention and MultiHeadAttention. The next component in the transformer block is the FeedForward network: a position-wise two-layer MLP that is applied independently to each token after attention.
- FeedForward network: A two-layer MLP with a hidden expansion factor (typically 4x). Takes each token's representation from
n_embdto4 * n_embdand back ton_embd, with a nonlinearity in between. - Layer normalization: Stabilizes activations before attention and before the feed-forward network. The pre-norm variant (used in GPT-2) applies LayerNorm before each sub-layer.
- Residual connections: Skip connections that add the input back to the output of each sub-layer:
x = x + Attention(LayerNorm(x))andx = x + FFN(LayerNorm(x)). - Full transformer block: Combining multi-head attention, feed-forward, layer norm, and residuals into one reusable block that can be stacked to build the complete model.