How to Implement Multi-Head Attention from Scratch in PyTorch
Implement multi-head attention in PyTorch by creating a Head class for scaled dot-product attention with causal masking, then wrapping multiple heads in a MultiHeadAttention class that concatenates their outputs to preserve input dimensions while enriching contextual representation.
To implement multi-head attention from scratch in PyTorch, you need to split the embedding dimension into parallel attention heads, compute scaled dot-product attention for each, and concatenate the results. The repository FareedKhan-dev/train-llm-from-scratch provides a minimal, educational implementation following the original "Attention Is All You Need" paper. The code demonstrates how to build the core Transformer mechanism without relying on high-level abstractions like nn.MultiheadAttention.
Understanding the Architecture
The implementation splits attention into two distinct classes to separate single-head logic from multi-head orchestration.
The Single Attention Head
The Head class in src/models/attention.py (lines 6-34) implements one attention head. It projects the input into queries, keys, and values using three separate nn.Linear layers. Each projection maps the full embedding dimension n_embed down to head_size = n_embed // n_head. Critically, these linear layers omit biases for simplicity and computational efficiency.
Combining Multiple Heads
The MultiHeadAttention class in src/models/attention.py (lines 60-96) manages parallelism. It initializes a nn.ModuleList of Head instances according to the n_head parameter. During the forward pass, it processes each head independently, concatenates their outputs along the embedding dimension (dim=-1), and returns a tensor with shape (B, T, n_embed) identical to the input, but containing richer contextual information from multiple representation subspaces.
Step-by-Step Implementation Details
The attention mechanism follows five precise operations as implemented in the source code:
-
Linear Projections – Three
nn.Linearlayers (key,query,value) transform the input fromn_embedtohead_sizedimensions. -
Scaled Dot-Product – Compute raw attention scores by multiplying queries (
q) with the transpose of keys (k.T), yielding a matrix of shape(B, T, T). Scale these scores by1/√head_size(stored asscale_factor) to prevent softmax saturation and maintain gradient stability. -
Causal Masking – Register a lower-triangular matrix (
self.tril) as a buffer during initialization. Apply this mask to the attention scores by setting future positions tofloat('-inf'), ensuring tokens only attend to previous positions for autoregressive generation. -
Softmax and Weighted Sum – Apply softmax over the last dimension to obtain attention weights, then multiply these weights with the values (
v) to produce the head output. -
Head Concatenation – Concatenate outputs from all heads along the embedding dimension, resulting in a tensor that maintains the original input shape but combines information from all attention heads.
Code Examples
Stand-Alone Multi-Head Attention
Use the MultiHeadAttention class directly for testing or integration into custom architectures. This example mirrors the validation block found at lines 98-111 of src/models/attention.py:
import torch
from src.models.attention import MultiHeadAttention
# Hyper-parameters
batch_size = 2
seq_len = 5
embed_dim = 32
num_heads = 4
context_len = seq_len
# Random input tensor (Batch, Time, Channels)
x = torch.randn(batch_size, seq_len, embed_dim)
# Initialize multi-head attention
mh_attn = MultiHeadAttention(
n_head=num_heads,
n_embed=embed_dim,
context_length=context_len,
)
# Forward pass
out = mh_attn(x)
print("Input shape :", x.shape) # → torch.Size([2, 5, 32])
print("Output shape:", out.shape) # → torch.Size([2, 5, 32])
Integration into Transformer Blocks
In practice, multi-head attention operates within a larger Transformer block. The Block class in src/models/transformer_block.py (lines 6-30) wraps the attention mechanism with layer normalization, residual connections, and a feedforward MLP:
from src.models.transformer_block import Block
block = Block(
n_head=num_heads,
n_embed=embed_dim,
context_length=context_len,
)
output = block(x) # Shape: (Batch, Time, embed_dim)
print(output.shape) # → torch.Size([2, 5, 32])
Source Code Structure
Reference these specific files in FareedKhan-dev/train-llm-from-scratch to explore the complete implementation:
-
src/models/attention.py– Contains theHeadandMultiHeadAttentionclass definitions implementing the core attention logic. -
src/models/transformer_block.py– Demonstrates how attention integrates with layer normalization, residual connections, and MLP layers to form a complete Transformer block. -
src/models/transformer.py– Stacks multipleBlockinstances to construct a full Transformer language model. -
scripts/train_transformer.py– Provides the training script that wires the model, dataset, and optimizer together for end-to-end training.
Summary
- Split embeddings across heads by setting
head_size = n_embed // n_headto enable parallel attention processing. - Scale attention scores by
1/√head_sizebefore softmax to maintain stable gradients during training. - Apply causal masking using a lower-triangular buffer to enforce autoregressive behavior in decoder-only models.
- Concatenate head outputs to preserve the original tensor shape while combining contextual information from multiple representation subspaces.
- Omit biases in projection layers for computational efficiency, following modern Transformer implementations.
Frequently Asked Questions
What is the difference between the Head and MultiHeadAttention classes?
The Head class implements a single attention mechanism with its own query, key, and value projections, computing scaled dot-product attention for one representation subspace. The MultiHeadAttention class instantiates multiple Head objects in parallel, processes the input through each independently, and concatenates their outputs to combine multiple perspectives simultaneously.
Why must attention scores be scaled by 1 divided by the square root of head size?
Scaling by 1/√head_size (also called the scale factor) prevents the dot-product values from growing too large as dimensionality increases. According to the Transformer paper, large dot-products push the softmax function into regions with extremely small gradients, which slows training convergence. The scaling factor keeps the variance of attention weights manageable and preserves gradient flow.
How does causal masking work in this implementation?
The implementation registers a lower-triangular boolean matrix (self.tril) as a buffer during initialization. During the forward pass, it masks the attention scores by setting future positions (above the diagonal) to negative infinity before the softmax operation. This ensures each position can only attend to itself and previous positions, which is essential for language modeling tasks where future tokens must remain hidden.
Can this implementation support bidirectional (non-causal) attention?
Yes, you can modify the implementation for bidirectional attention by removing or conditionally applying the causal mask. In src/models/attention.py, you would either skip the masking step where self.tril is applied, or add a boolean parameter (e.g., causal=True) that controls whether to apply the mask. This flexibility allows the same code to support both decoder-only (causal) and encoder-style (bidirectional) Transformer architectures.
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 →