How Multi-Head Attention Works in PyTorch: A Deep Dive into the Transformer Implementation
Multi-head attention enables a Transformer model to relate every token to every other token by computing scaled dot-product attention in parallel across multiple representation subspaces, then concatenating and linearly projecting the results to produce context-aware embeddings.
Multi-head attention forms the backbone of modern language models, allowing them to weigh the importance of different tokens when encoding context. In the FareedKhan-dev/train-llm-from-scratch repository, this mechanism is built from three tightly coupled PyTorch components that process input embeddings through parallel attention heads and causal masking.
The Building Blocks of Multi-Head Attention
Single Attention Head (Head class)
In src/models/attention.py, the Head class implements a single attention mechanism. It first projects input embeddings into key, query, and value matrices using learned linear transformations. The attention scores are computed as the scaled dot-product Q @ K^T / sqrt(head_size), where the scaling factor stabilizes gradients. A causal mask (implemented via torch.tril) is applied to prevent the model from attending to future tokens, ensuring autoregressive behavior essential for language modeling.
Parallel Processing with MultiHeadAttention
The MultiHeadAttention class orchestrates n_head parallel instances of Head. Each head receives the same layer-normalized input but processes it through independent key, query, and value projections. The outputs from all heads are concatenated along the embedding dimension (dim=-1) and passed through a final linear projection (proj) to blend the information. This architecture preserves the input tensor shape (B, T, C)—where B is batch size, T is sequence length, and C is embedding dimension—while enriching representations with context from multiple subspaces.
The Complete Transformer Block Flow
The Block class in src/models/transformer_block.py integrates multi-head attention into a standard Transformer architecture:
- LayerNorm:
ln1(x)normalizes the input embeddings - Multi-Head Attention:
headscompute parallel attention, are concatenated, and projected - Residual Connection:
x + attn_outadds the attention output to the original input - LayerNorm:
ln2(x)normalizes the result before the MLP - MLP: A feed-forward network applies non-linear transformations
- Residual Connection:
x + mlp_outproduces the final block output
The causal mask guarantees that each token only attends to previous tokens, maintaining the temporal order required for next-token prediction.
Why Multiple Heads Matter
Each attention head learns to focus on different aspects of the embedding space. By distributing computation across multiple heads, the model captures diverse linguistic relationships—such as syntactic dependencies and long-range semantic associations—without increasing the dimensionality of individual heads. The final projection layer (proj) synthesizes these parallel perspectives into a unified representation.
Implementation Examples
Standalone MultiHeadAttention
import torch
from src.models.attention import MultiHeadAttention
# Parameters
batch_size = 2
seq_len = 8
embed_dim = 32
num_heads = 4
# Random input (B, T, C)
x = torch.randn(batch_size, seq_len, embed_dim)
# Multi-head attention module
mh_attn = MultiHeadAttention(n_head=num_heads, n_embed=embed_dim, context_length=seq_len)
# Forward pass
out = mh_attn(x)
print("Input shape :", x.shape) # (2, 8, 32)
print("Output shape:", out.shape) # (2, 8, 32)
Integrated Transformer Block
import torch
from src.models.transformer_block import Block
batch_size = 2
seq_len = 8
embed_dim = 32
num_heads = 4
x = torch.randn(batch_size, seq_len, embed_dim)
# Complete transformer block with attention, layer norm, and MLP
block = Block(n_head=num_heads, n_embed=embed_dim, context_length=seq_len)
y = block(x)
print("Block output shape:", y.shape) # (2, 8, 32)
Key Source Files and Architecture
src/models/attention.py: ImplementsHeadandMultiHeadAttentionclasses with causal masking and scaled dot-product attentionsrc/models/transformer_block.py: Wraps multi-head attention with layer normalization, residual connections, and MLP layers to form a complete Transformer blocksrc/models/transformer.py: Stacks multipleBlockinstances, manages token embeddings, and provides inference utilities
Summary
- Multi-head attention splits computation across parallel heads to capture diverse token relationships in different embedding subspaces
- The
Headclass insrc/models/attention.pyimplements scaled dot-product attention with a causal mask (tril) to prevent future-token leakage MultiHeadAttentionconcatenates outputs fromn_headinstances and projects them back to the embedding dimension using a learned linear layer- Transformer blocks combine attention with MLPs using layer normalization (
ln1,ln2) and residual connections to enable deep network training - The implementation maintains tensor shapes
(B, T, C)throughout the forward pass, ensuring dimensional consistency across the model architecture
Frequently Asked Questions
What is the difference between single-head and multi-head attention?
Single-head attention uses one set of query, key, and value projections to compute attention scores. Multi-head attention instantiates n_head parallel heads, each learning different aspects of the data, then concatenates their outputs. This allows the model to jointly attend to information from different representation subspaces at different positions, capturing richer contextual relationships than a single head could achieve.
Why is the causal mask necessary in the attention mechanism?
The causal mask (implemented as a lower-triangular matrix via torch.tril) enforces autoregressive generation by preventing tokens from attending to future positions in the sequence. During training, this ensures that the prediction for position i only depends on tokens at positions 1 through i, maintaining the temporal order required for language modeling and preventing information leakage from target tokens.
How does the scaling factor 1/sqrt(head_size) improve training stability?
Without scaling, the dot-product scores grow with the dimensionality of the key vectors, potentially pushing softmax outputs into regions with extremely small gradients. Dividing by the square root of the head dimension keeps the variance of attention scores approximately constant regardless of embedding size, enabling more stable gradient flow and faster convergence during training.
What shape should the input tensor have when using MultiHeadAttention?
The input tensor must have shape (B, T, C) where B is the batch size, T is the sequence length (context length), and C is the embedding dimension (n_embed). The MultiHeadAttention module preserves this shape throughout the forward pass, returning a tensor of identical dimensions (B, T, C) after concatenating head outputs and applying the projection layer.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →