← All Posts

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:

  1. Runs multiple SingleHeadAttention modules in parallel. Each head receives the same input tensor but processes it through its own learned Q, K, V projections.
  2. 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.
  3. 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.
The original paper's formulation: Vaswani et al. define:
MultiHead(Q, K, V) = Concat(head₁, …, headₕ) · W_O where each headᵢ = Attention(X · W_Qᵢ, X · W_Kᵢ, X · W_Vᵢ)

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:

Input: (batch_size, seq_len, n_embd) Output: (batch_size, seq_len, n_embd)

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:

class MultiHeadAttention : public torch::nn::Module{ /* It has access to multiple singlheadattention, and then it concatenates the output of all of these layer into one vector which then it passes to a projection layer input : batch_size x seq_len x n_embd output : batch_size x seq_len x (#of heads x head_size => n_embd) */ int num_heads, head_size, n_embd; vector< shared_ptr<SingleHeadAttention> > SHA_group; torch::nn::Linear projection_layer{nullptr}; public: MultiHeadAttention(int num_heads, int head_size, int n_embd) : num_heads(num_heads), head_size(head_size), n_embd(n_embd){ for (int head_idx = 0; head_idx < num_heads; head_idx++){ shared_ptr<SingleHeadAttention> curr_head = make_shared<SingleHeadAttention>(n_embd, head_size); register_module( "SHA_" + to_string(head_idx), curr_head ); SHA_group.push_back(curr_head); } projection_layer = register_module( "projection_layer", torch::nn::Linear( torch::nn::LinearOptions( head_size * num_heads, n_embd ) ) ); } torch::Tensor forward(torch::Tensor x){ torch::Tensor output, concatenated_SHA_outputs; vector < torch::Tensor> SHA_outputs; for(int head_idx = 0; head_idx < num_heads; head_idx++){ SHA_outputs.push_back(SHA_group[head_idx]->forward(x)); } concatenated_SHA_outputs = torch::cat(SHA_outputs, -1); output = projection_layer(concatenated_SHA_outputs); return output; } };

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:

for (int head_idx = 0; head_idx < num_heads; head_idx++){ shared_ptr<SingleHeadAttention> curr_head = make_shared<SingleHeadAttention>(n_embd, head_size); register_module( "SHA_" + to_string(head_idx), curr_head ); SHA_group.push_back(curr_head); }

Three things happen on each iteration:

  1. 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).
  2. 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.
  3. Store the head: SHA_group.push_back(curr_head) keeps a reference in the vector for use during the forward pass.
Why register_module inside the loop? Each head is a separate torch::nn::Module with its own parameters. If you skip registration, those parameters become invisible to the optimizer and to model->to(torch::kCUDA).
Note on 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

projection_layer = register_module( "projection_layer", torch::nn::Linear( torch::nn::LinearOptions( head_size * num_heads, n_embd ) ) );

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.

n_embd = 256, num_heads = 8 head_size = n_embd / num_heads = 256 / 8 = 32 Each head output: (B, S, 32) Concatenated: (B, S, 32 * 8) = (B, S, 256) = (B, S, n_embd)

Some more examples to build intuition:

n_embdnum_headshead_sizeEach head outputConcatenated
256464(B, S, 64)(B, S, 256)
256832(B, S, 32)(B, S, 256)
512864(B, S, 64)(B, S, 512)
7681264(B, S, 64)(B, S, 768)
10241664(B, S, 64)(B, S, 1024)
n_embd must be divisible by num_heads. If 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.
Important distinction: Each head does not receive a slice of the input. Every head receives the complete 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:

torch::Tensor forward(torch::Tensor x){ torch::Tensor output, concatenated_SHA_outputs; vector < torch::Tensor> SHA_outputs; for(int head_idx = 0; head_idx < num_heads; head_idx++){ SHA_outputs.push_back(SHA_group[head_idx]->forward(x)); } concatenated_SHA_outputs = torch::cat(SHA_outputs, -1); output = projection_layer(concatenated_SHA_outputs); return output; }

Let us walk through every line.

Line 1: The head loop

for(int head_idx = 0; head_idx < num_heads; head_idx++){ SHA_outputs.push_back(SHA_group[head_idx]->forward(x)); }

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

concatenated_SHA_outputs = torch::cat(SHA_outputs, -1);

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).

SHA_outputs[0]: (B, S, 32) SHA_outputs[1]: (B, S, 32) ... SHA_outputs[7]: (B, S, 32) torch::cat(SHA_outputs, -1) → (B, S, 256)

Line 3: Projection

output = projection_layer(concatenated_SHA_outputs);

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

return output;

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

Input x (1, 4, 96) Head 0 (SHA_0) Q, K, V, mask, softmax (1, 4, 32) Head 1 (SHA_1) Q, K, V, mask, softmax (1, 4, 32) Head 2 (SHA_2) Q, K, V, mask, softmax (1, 4, 32) 32 dims 32 dims 32 dims torch::cat(dim=-1) (1, 4, 96) projection_layer Linear(96, 96) Output (1, 4, 96)
Click Step to begin

6 · Why shared_ptr and Not unique_ptr?

LibTorch's register_module requires a shared_ptr. When you call register_module("SHA_0", curr_head), the module registry stores a copy of the pointer. Now two things reference the head: your SHA_group vector and the registry. A unique_ptr cannot be shared by definition — it would need to be moved, leaving your vector with a null pointer.

Why LibTorch chose shared ownership internally: a registered submodule is accessed through multiple code paths — parameters() walks the registry to collect all learnable tensors for the optimizer, to(device) traverses it to move weights to GPU, save/load serialises from it, and your own forward() accesses the same object through the vector. All of these paths must hold a valid reference to the same underlying module simultaneously. Shared ownership via shared_ptr is the simplest guarantee: the module stays alive as long as any owner still holds a reference, and the reference count handles cleanup automatically.

This also means copies of the parent module (e.g., cloning for data parallelism) get their own reference to each submodule without deep-copying weights until explicitly told to.

// This will NOT compile: unique_ptr<SingleHeadAttention> head = make_unique<SingleHeadAttention>(n_embd, head_size); register_module("SHA_0", head); // ERROR: no matching overload
Rule of thumb: Any torch::nn::Module subclass registered as a submodule must be managed with shared_ptr. Use make_shared to create it, then pass to both register_module and your own storage.

7 · Parameter Count

Using concrete numbers: n_embd = 256, num_heads = 8, head_size = 32.

Per head: 3 * (n_embd * head_size) = 3 * 8,192 = 24,576 (Q + K + V, no bias) All heads: 8 * 24,576 = 196,608 Projection: n_embd * n_embd + n_embd = 65,536 + 256 = 65,792 (weight + bias) ------- Grand total: 262,400

The general formula simplifies to 4 * n_embd² + n_embd, which is independent of the number of heads.

Modeln_embdnum_headshead_sizeMHA Params
GPT-2 Small76812642,360,064
GPT-2 Medium102416644,195,328
GPT-2 Large128020646,554,880
GPT-2 XL1600256410,241,600
Note: These counts are for a single MHA layer. A full model stacks many layers (12–48), so total attention parameters are multiplied accordingly.

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_embd to 4 * n_embd and back to n_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)) and x = 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.
The pattern so far: SingleHeadAttention learns one attention pattern in a subspace. MultiHeadAttention runs several in parallel, concatenates, and projects. Next, the FeedForward network provides the non-linear transformation capacity that attention alone cannot provide. Together, attention plus feed-forward form the backbone of every transformer layer.