How to Implement Self-Attention Mechanism from Scratch in Python
Self-attention enables transformer models to relate every token in a sequence to every other token by computing scaled dot-product attention between learned query, key, and value projections, implemented efficiently in PyTorch using a single combined QKV linear layer and causal masking for autoregressive generation.
Self-attention forms the computational backbone of modern large language models, allowing sequences to capture long-range dependencies without recurrence. The rohitg00/ai-engineering-from-scratch repository provides a complete, educational implementation that demonstrates exactly how to build this mechanism using pure PyTorch operations. This guide walks through the actual source code from the capstone projects, covering everything from linear projections to multi-head processing.
The Seven Steps of Self-Attention
The MultiHeadSelfAttention class in phases/19-capstone-projects/33-multihead-self-attention/code/main.py encapsulates the complete attention mechanism through seven distinct operations:
- Linear projections – A single learned layer generates queries (Q), keys (K), and values (V) from the input representation
xwith shapeB×T×D - Head splitting – The concatenated QKV tensors reshape so each head processes a sub-space of dimension
d_head = D / n_heads - Scaled dot-product – Attention scores compute as
Q·Kᵀscaled by√d_headto maintain stable softmax gradients - Causal mask – A lower-triangular mask prevents tokens from attending to future positions in autoregressive setups
- Softmax and dropout – Masked scores convert to probabilities via
softmax, followed by optional regularization - Weighted sum – The probability distribution combines values through matrix multiplication
weights·V - Head merging and output projection – Per-head results concatenate and project back to the original dimension
Building the MultiHeadSelfAttention Class
Combined QKV Projection
For computational efficiency, the implementation uses a single linear layer to project the input into query, key, and value representations simultaneously. In phases/19-capstone-projects/33-multihead-self-attention/code/main.py at line 47, the initialization creates:
self.qkv_proj = nn.Linear(d_model, 3 * d_model, bias=True)
This concatenated projection reduces memory overhead and kernel launch costs compared to three separate layers. The forward pass splits the output tensor into three equal parts corresponding to Q, K, and V.
Splitting and Merging Heads
Multi-head attention divides the model dimension across parallel attention heads. The _split_heads method reshapes the tensor using view and transpose operations:
x.view(b, t, n_heads, d_head).transpose(1, 2)
This transforms the shape from (batch, seq, d_model) to (batch, n_heads, seq, d_head), allowing each head to attend to the full sequence independently. After computing attention, the _merge_heads method reverses this operation at line 60:
x.transpose(1, 2).contiguous().view(b, t, h * dh)
Scaled Dot-Product Attention
The core attention computation occurs at line 83, where the implementation calculates compatibility scores between queries and keys:
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_head)
Scaling by the square root of the head dimension prevents the dot products from growing too large in magnitude, which would push the softmax function into regions with extremely small gradients.
Causal Masking Implementation
For language modeling tasks, the class maintains a causal mask buffer initialized at lines 52-53:
self.causal_mask = torch.tril(torch.ones(max_context_length, max_context_length))
During the forward pass, the implementation slices this mask to the current sequence length and applies it using masked_fill:
scores = scores.masked_fill(mask_slice == 0, float("-inf"))
This ensures positions can only attend to themselves and previous tokens, maintaining the autoregressive property required for next-token prediction.
Practical Implementation Examples
Minimal Self-Attention Demo
To use the implementation directly for experimentation:
import torch
from phases.19_capstone_projects.33_multihead_self_attention.code.main import MultiHeadSelfAttention
# Configuration
batch, seq_len, d_model = 2, 5, 16
n_heads = 4
# Random input
x = torch.randn(batch, seq_len, d_model)
# Initialise the layer
attn = MultiHeadSelfAttention(d_model=d_model,
n_heads=n_heads,
max_context_length=seq_len)
# Forward pass – returns (output, attention_weights)
out, weights = attn(x, return_weights=True)
print("output shape :", out.shape) # → (2, 5, 16)
print("weights shape:", weights.shape) # → (2, 4, 5, 5)
Tiny Language Model Integration
The repository includes TinyAttentionLM, which demonstrates how self-attention functions within a complete architecture:
from phases.19_capstone_projects.33_multihead_self_attention.code.main import TinyAttentionLM, DemoConfig
cfg = DemoConfig(vocab_size=64, d_model=32, n_heads=4, seq_len=12, batch_size=2)
model = TinyAttentionLM(vocab_size=cfg.vocab_size,
d_model=cfg.d_model,
n_heads=cfg.n_heads,
max_context_length=cfg.seq_len)
# Dummy token ids
ids = torch.randint(0, cfg.vocab_size, (cfg.batch_size, cfg.seq_len))
# Forward pass – get logits and per‑head attention weights
logits, attn_weights = model(ids, return_weights=True)
print("logits shape :", logits.shape) # → (2, 12, 64)
print("attention weights shape :", attn_weights.shape) # → (2, 4, 12, 12)
Full Transformer Block Stack
For production-like usage, the attention mechanism integrates with LayerNorm and feed-forward networks in phases/19-capstone-projects/34-transformer-block/code/main.py:
from phases.19_capstone_projects.34_transformer_block.code.main import BlockConfig, BlockStack
cfg = BlockConfig(d_model=128, num_heads=8, context_length=32, pre_ln=True)
stack = BlockStack(cfg, depth=3) # three transformer blocks
tokens = torch.randint(0, 128, (2, 32)) # batch of token ids
output = stack(tokens)
print("stack output shape:", output.shape) # → (2, 32, 128)
Summary
- Single QKV projection improves efficiency by computing queries, keys, and values in one matrix multiplication using
nn.Linear(d_model, 3 * d_model) - Head splitting via
viewandtransposeoperations enables parallel attention across multiple representation subspaces without increasing computational complexity - Causal masking uses a pre-allocated lower-triangular buffer to prevent attention to future tokens during training
- Scaling factor of
1/√d_headmaintains stable gradient flow through the softmax operation - Modular design allows the
MultiHeadSelfAttentionclass to function standalone or within complete transformer blocks containing residual connections and layer normalization
Frequently Asked Questions
Why combine QKV into a single linear projection?
Combining the query, key, and value projections into one layer reduces memory overhead and enables optimized kernel fusion during training. The implementation splits the output tensor into three parts, achieving identical mathematical results to separate layers while improving computational efficiency on modern hardware accelerators.
What is the purpose of scaling by the square root of the head dimension?
Scaling by √d_head prevents the dot-product values from becoming too large as the dimension increases. Without this scaling, the softmax function would receive extremely large inputs, producing gradients near zero that slow down or prevent effective learning.
How does causal masking prevent looking at future tokens?
The causal mask is a lower-triangular matrix where positions (i, j) contain 1 if j ≤ i and 0 otherwise. Before applying softmax, the implementation sets masked positions to negative infinity using masked_fill, causing the softmax output to be zero for future positions. This ensures each token can only aggregate information from itself and previous tokens in the sequence.
Can this implementation handle bidirectional attention?
Yes, by setting the causal mask parameter to None or modifying the mask initialization logic, the same MultiHeadSelfAttention class supports bidirectional attention for encoder-only architectures like BERT. The mask buffer is optional, and the forward pass can skip the masking step when processing non-autoregressive tasks.
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 →