How to Implement Multi-Head Attention from Scratch Using PyTorch
Multi-head attention splits the embedding dimension into parallel heads that each compute scaled dot-product attention with causal masking; the train-llm-from-scratch repository implements this via a Head class for single-head logic and a MultiHeadAttention class that concatenates multiple heads along the embedding dimension to preserve the input shape (B, T, n_embed).
Multi-head attention is the core mechanism powering modern Transformer architectures. This article examines a minimal PyTorch implementation from the FareedKhan-dev/train-llm-from-scratch repository, demonstrating exactly how to implement multi-head attention from scratch using PyTorch without relying on high-level abstractions.
Understanding the Multi-Head Attention Architecture
The implementation follows the original "Attention Is All You Need" design, separating concerns into two distinct classes within src/models/attention.py. The Head class manages single-head attention logic, while MultiHeadAttention orchestrates parallel computation across multiple heads.
This modular approach allows each head to learn different representation subspaces while maintaining a clean interface for the broader Transformer model.
Step-by-Step Implementation
Single Attention Head (Head class)
Located at lines 6-34 in src/models/attention.py, the Head class implements scaled dot-product attention with three critical operations:
Linear projections. For each input tensor of shape (B, T, n_embed), three nn.Linear layers—key, query, and value—project the embedding into a lower-dimensional head_size. The implementation omits biases for computational efficiency, mapping to head_size = n_embed // n_head.
Scaled dot-product computation. The attention scores are computed as the matrix multiplication of queries (q) and transposed keys (k.transpose(-2, -1)), producing a raw score matrix of shape (B, T, T). These scores are scaled by 1 / sqrt(head_size) (the scale_factor) to prevent softmax saturation and stabilize gradients.
Causal masking. A lower-triangular matrix registered as buffer self.tril masks future positions with float('-inf'), enforcing autoregressive behavior where each token attends only to previous tokens.
The final output results from applying softmax to the masked scores and multiplying by the values tensor (v).
Parallel Head Aggregation (MultiHeadAttention class)
The MultiHeadAttention class (lines 60-96 in src/models/attention.py) aggregates multiple Head instances:
-
ModuleList initialization. The constructor creates
n_headinstances ofHead, each processing a slice of the embedding dimension in parallel. -
Concatenation strategy. During the forward pass, each head produces an output of shape
(B, T, head_size). The implementation concatenates these along the last dimension (dim=-1), yielding a tensor of shape(B, T, n_embed)that matches the input dimensions but contains richer contextual information.
This shape preservation allows the attention output to be directly added to the residual stream within Transformer blocks.
Complete Code Examples
Stand-Alone Multi-Head Attention
The following example mirrors the test block at lines 98-111 in src/models/attention.py, demonstrating initialization and forward pass:
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
# 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 preserves input shape
out = mh_attn(x)
print(f"Input shape: {x.shape}") # torch.Size([2, 5, 32])
print(f"Output shape: {out.shape}") # torch.Size([2, 5, 32])
Integration into a Transformer Block
In practice, the attention module integrates into a full Transformer block via src/models/transformer_block.py (lines 6-30), which wraps the attention with layer normalization and residual connections:
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 (B, T, embed_dim)
print(output.shape) # torch.Size([2, 5, 32])
Key Implementation Details
Several design choices in FareedKhan-dev/train-llm-from-scratch optimize the implementation for clarity and training stability:
-
Head dimension calculation. The
head_sizeis strictly derived asn_embed // n_head, ensuring the concatenated heads exactly reconstruct the original embedding dimension without additional projection layers. -
Buffer-based masking. The causal mask
self.trilis registered as a persistent buffer rather than computed dynamically, reducing overhead during training loops. -
Biases disabled. All
nn.Linearprojections for keys, queries, and values setbias=False, following modern LLM implementations to reduce parameter count without impacting representational capacity.
Summary
- The
Headclass insrc/models/attention.pyimplements single-head scaled dot-product attention with causal masking via a lower-triangular buffer. MultiHeadAttentionparallelizesn_headinstances usingnn.ModuleListand concatenates results alongdim=-1to preserve the(B, T, n_embed)tensor shape.- Scaled attention (
1/sqrt(head_size)) prevents softmax gradient vanishing in deep models. - The implementation integrates into full Transformer blocks via
src/models/transformer_block.py, supporting autoregressive language modeling workflows.
Frequently Asked Questions
Why is the attention score scaled by the square root of the head size?
Scaling by 1 / sqrt(head_size) prevents the dot-product values from growing too large in magnitude as the dimensionality increases. Without this scaling, the softmax function would produce extremely small gradients, slowing or destabilizing training. This scaling factor appears in lines 6-34 of src/models/attention.py.
What is the purpose of the causal mask in the Head class?
The causal mask ensures autoregressive behavior by preventing tokens from attending to future positions in the sequence. Implemented via the self.tril buffer in src/models/attention.py, it masks the upper triangle of the attention score matrix with float('-inf') before the softmax operation, effectively setting future attention weights to zero.
How does MultiHeadAttention maintain the input-output shape consistency?
The class divides the embedding dimension evenly across heads (head_size = n_embed // n_head), processes each head in parallel, then concatenates the results along the last dimension. Because n_head * head_size = n_embed, the output tensor retains the original shape (B, T, n_embed), allowing seamless residual connections in the Transformer architecture.
Can this implementation handle variable sequence lengths?
Yes. The implementation uses PyTorch tensor operations that support dynamic shapes. While the context_length parameter initializes the causal mask buffer, the forward pass in src/models/attention.py handles any sequence length up to this maximum, making it suitable for training and inference with varying text lengths.
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 →