Hello! Welcome back to our deep dive into the mechanisms that power modern AI.
In our last lesson, we derived and implemented scaled dot-product self-attention. We established the Query, Key, and Value (Q, K, V) framework and understood the critical role of scaling to ensure stable training. This gave us a powerful tool for a sequence to "look" at itself and build context-rich representations.
However, a single attention mechanism might be forced to average its focus, potentially missing out on diverse and subtle relationships within the text. For example, in the sentence "The tired robot dropped the heavy gear because it broke," we might want to simultaneously understand that "it" refers to "gear" (co-reference) and that "broke" is linked to "robot dropped" (causal relationship). A single attention head might struggle to capture both.
This brings us to today's topic: Multi-Head Attention. Your learning outcome for this lesson is to build a multi-head attention layer, the core component of the Transformer. We will see how running multiple attention mechanisms in parallel allows the model to capture a richer variety of linguistic features and dependencies.
1. The Rationale: From Single Head to Many
The core idea behind multi-head attention is to run the scaled dot-product attention process multiple times in parallel with different, learned linear projections. This allows the model to jointly attend to information from different "representation subspaces" at different positions.
You can think of it as giving the model multiple "perspectives" on the input sequence. Each "head" can learn to focus on a different type of relationship:
- One head might track syntactic dependencies, like subject-verb agreement.
- Another might track co-reference, like which pronoun refers to which noun.
- A third might focus on positional information, like which words are nearby.
By combining the outputs of these heads, the model can construct a far more nuanced and comprehensive representation of each word in the sequence.
2. The Architecture of Multi-Head Attention
Let's break down the architecture. At a high level, the process is:
- Project: For each of the
hheads, create separate, learned weight matrices to project the input Queries, Keys, and Values into a lower-dimensional space. - Attend in Parallel: Perform scaled dot-product attention for each head simultaneously, using its projected Q, K, and V. This yields
hseparate output matrices. - Concatenate & Combine: Concatenate the
houtput matrices back together and pass them through a final, learned linear projection to produce the layer's final output.
This architecture is elegantly captured in the following diagram.

The mathematical formulation from the original "Attention Is All You Need" paper is as follows:
Let's dissect this:
- The inputs are the full-dimensional
Q,K,Vmatrices (e.g., dimensiond_model= 512). - For each head
i, we have unique weight matrices . These project the inputs down to the head's dimension (d_k, e.g., 64). Attention(...)is the scaled dot-product attention function we learned in the last lesson.- The results from all heads are concatenated.
- A final weight matrix, , is applied to combine the information from all heads and project the result back to the original
d_modeldimension.
To build a strong visual intuition for this process, the following resource is excellent.
Jay Alammar's 'The Illustrated Transformer' provides one of the clearest explanations of multi-head attention. Pay close attention to the diagrams showing the multiple sets of Q/K/V matrices and how their outputs are ultimately combined.
Read the section titled 'The Beast With Many Heads'. It walks through the process of creating multiple Q/K/V sets, performing attention for each, and then concatenating and projecting the results. The visuals make the flow of matrices very clear.
3. Implementation in PyTorch
Now, let's translate this architecture into code. Your background in Python and software engineering will be valuable here as we construct a MultiHeadAttention module. The key challenge lies in efficiently managing the tensor shapes as we split into heads and then merge back.
The most efficient implementation doesn't create h separate linear layers. Instead, it uses single, larger linear layers and then reshapes the tensors. The process for the forward pass is:
- Define linear layers to map the input
d_modeltod_modelfor Q, K, and V. - Pass the input
query,key, andvaluetensors through these layers. - Reshape the resulting
(batch_size, seq_len, d_model)tensors into(batch_size, seq_len, num_heads, head_dim). - Permute the dimensions to
(batch_size, num_heads, seq_len, head_dim)to facilitate batch matrix multiplication across all heads at once. - Apply the scaled dot-product attention formula using the reshaped tensors.
- Permute and reshape the output back to
(batch_size, seq_len, d_model). - Pass this through a final linear output layer.
The following video provides a complete, line-by-line implementation of this process. It's an excellent resource for seeing exactly how the tensor manipulations work in PyTorch.
Multi Head Attention in Transformer Neural Networks with Code!
This video from CodeEmporium is focused exclusively on coding the multi-head attention mechanism. It explains the tensor reshaping and calculations step-by-step before encapsulating it all in a PyTorch class.
Watch the entire video (about 15 minutes). It is a complete walkthrough that perfectly matches our goal. (02:42): Note the conceptual breakdown into Q, K, V and then into multiple heads. (06:25): Focus on the reshaping of the concatenated QKV tensor to accommodate the heads. The permutation qkv.permute(0, 2, 1, 3) is key to preparing for parallel computation. (07:30): This is where the core scaled dot-product attention calculation happens. You'll recognize the steps from our previous lesson. (11:37): See how the outputs from the heads are concatenated and passed through a final linear layer. (13:42): The final MultiHeadedAttention class brings all these pieces together.
While the CodeEmporium video uses a combined linear layer for Q, K, and V, an alternative and very common approach is to use three separate linear layers. Aladdin Persson's video on building a Transformer from scratch demonstrates this, and also introduces the torch.einsum function—a powerful tool for expressing complex tensor operations that might appeal to your programming background.
Pytorch Transformers from Scratch (Attention is all you need)
For a different perspective and to see how this layer fits into the bigger picture, this video by Aladdin Persson is invaluable. He implements the same logic but calls the class SelfAttention and uses some different PyTorch techniques.
Watch from 11:28 to 27:05. Focus on his SelfAttention class implementation. Notice how he defines separate linear layers for values, keys, and queries within the __init__ method, but applies them later. His use of torch.einsum for the matrix multiplications is a very concise and powerful alternative to torch.matmul and explicit transposing. Try to understand how the string notation ('nqhd,nkhd->nhqk') defines the multiplication and output shape.
Test your understanding!
Let's say you're building a multi-head attention layer with d_model = 512 and num_heads = 8. An input tensor x with shape (batch_size=32, seq_len=60, embed_dim=512) is passed in as the query, key, and value.
- What is the dimension of each head (
head_dim)? - After passing
xthrough the initial linear layer for the queries and reshaping/permuting it, what will be the shape of the resultingquerytensor ready for the attention calculation? - During the attention calculation
torch.matmul(query, key.transpose(-2, -1)), what is the shape of the resultingenergytensor?
Show answer
head_dim: It'sd_model / num_heads= 512 / 8 = 64.querytensor shape:- Initial projection:
(32, 60, 512) - Reshape to add heads:
(32, 60, 8, 64) - Permute for batch matmul:
(32, 8, 60, 64). The shape is(batch_size, num_heads, seq_len, head_dim).
- Initial projection:
energytensor shape:queryshape:(32, 8, 60, 64)key.transpose(-2, -1)shape:(32, 8, 64, 60)- The matrix multiplication is performed on the last two dimensions (
60x64@64x60), resulting in a shape of(32, 8, 60, 60). This is the attention score matrix for each of the 8 heads and each of the 32 examples in the batch.
Conclusion
Congratulations! You have now constructed the single most important component of the Transformer architecture. By understanding how to split the model's representation into parallel "heads," you've unlocked the mechanism that allows Transformers to capture rich and diverse relationships within data.
Key Takeaways:
- Multi-Head Attention improves on single-head attention by allowing the model to focus on different information from different representational subspaces simultaneously.
- The implementation involves projecting the Q, K, V inputs for each head, running scaled dot-product attention in parallel, and then concatenating and projecting the results.
- Efficient implementation relies on clever tensor reshaping and permutation to perform calculations for all heads at once using batched matrix multiplications.
Preview of the Next Lesson:
We are now ready to assemble a full layer of the Transformer. In the next lesson, we will take the MultiHeadAttention layer we've just built and wrap it with the other essential components: residual connections and layer normalization, followed by a position-wise feed-forward network. This will complete our construction of a Transformer encoder block.