Implementing Scaled Dot-Product Attention with Key, Query, and Value Projections
Implement key, query, and value projections using bias-free nn.Linear layers, compute scaled dot-product attention scores with causal masking, and aggregate outputs through weighted sums to build multi-head transformer attention mechanisms.
The train-llm-from-scratch repository provides a pure PyTorch implementation of transformer attention mechanisms without relying on high-level abstractions. Learning how to implement attention with key query value projections from scratch reveals the mathematical foundations that power modern large language models.
Projecting Inputs into Key, Query, and Value Representations
The attention mechanism begins by transforming input embeddings into three distinct representations. In [src/models/attention.py](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L29-L31), the Head class initializes three separate nn.Linear projection layers:
self.key = nn.Linear(n_embed, head_size, bias=False)
self.query = nn.Linear(n_embed, head_size, bias=False)
self.value = nn.Linear(n_embed, head_size, bias=False)
Bias is intentionally disabled in these projections because the subsequent scaled dot-product calculation operates on relative similarities; additive constants would cancel out during the matrix multiplication and softmax operations. Each layer maps the input embedding dimension (n_embed) to the head-specific dimension (head_size).
During the forward pass, the input tensor x of shape (B, T, n_embed)—where B represents batch size and T represents sequence length—undergoes simultaneous projection:
k = self.key(x) # (B, T, head_size)
q = self.query(x) # (B, T, head_size)
v = self.value(x) # (B, T, head_size)
Computing Scaled Dot-Product Attention
After projection, the implementation calculates attention scores using scaled dot-product similarity. According to the source code in [src/models/attention.py](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/attention.py#L49-L51), the operation follows the standard transformer formula:
attn_scores = q @ k.transpose(-2, -1) * (1.0 / math.sqrt(k.size(-1)))
Scaling by the square root of the head size prevents the dot-product values from growing too large in high-dimensional spaces, which would push the softmax function into regions with extremely small gradients.
To maintain autoregressive properties required for language modeling, the repository applies a causal mask using torch.tril (lines 52-54). This triangular mask ensures tokens only attend to themselves and previous positions, never to future tokens.
Aggregating Values and Building Multi-Head Output
The attention weights undergo softmax normalization to create a probability distribution (line 55), followed by a weighted aggregation with the value vectors (line 56):
attn_weights = F.softmax(attn_scores, dim=-1)
out = attn_weights @ v # (B, T, head_size)
This computation occurs independently across multiple heads. The MultiHeadAttention class (line 71) instantiates n_head separate Head modules, processing the same input through parallel projections. The outputs are concatenated along the feature dimension and passed through a final linear projection (line 78):
class MultiHeadAttention(nn.Module):
def __init__(self, n_head, n_embed, head_size, dropout):
super().__init__()
self.heads = nn.ModuleList([Head(n_embed, head_size, dropout) for _ in range(n_head)])
self.proj = nn.Linear(n_embed, n_embed) # Line 78
def forward(self, x):
out = torch.cat([h(x) for h in self.heads], dim=-1)
return self.proj(out)
Integrating Attention into Transformer Blocks
The attention mechanism serves as the core component within larger transformer architectures. In [src/models/transformer_block.py](https://github.com/FareedKhan-dev/train-llm-from-scratch/blob/main/src/models/transformer_block.py), the MultiHeadAttention module is combined with layer normalization, residual connections, and feed-forward networks:
from src.models.transformer_block import TransformerBlock
import torch
# Configuration
batch_size, seq_len, embed_dim = 2, 10, 64
n_head = 4
head_size = embed_dim // n_head
# Initialize transformer block
block = TransformerBlock(
n_head=n_head,
n_embed=embed_dim,
head_size=head_size,
dropout=0.1
)
# Forward pass
x = torch.randn(batch_size, seq_len, embed_dim)
output = block(x) # Shape: (2, 10, 64)
Summary
- Key, query, and value projections are implemented as separate
nn.Linearlayers without bias in theHeadclass to preserve relative similarity calculations. - Scaled dot-product attention divides scores by √head_size to maintain stable gradients during training.
- Causal masking via
torch.trilenforces autoregressive behavior by preventing future token visibility. - Multi-head concatenation combines parallel attention computations through a final linear projection at line 78.
- The complete architecture stacks these attention blocks within
TransformerBlockmodules to build full transformer models.
Frequently Asked Questions
Why are biases disabled in the key, query, and value projections?
Biases are unnecessary in attention projections because the scaled dot-product operation (Q @ K^T) computes relative similarities between vectors. Any additive constant introduced by bias would appear in both the query and key calculations, effectively canceling out during the dot-product operation while adding redundant parameters to the model.
What purpose does the causal mask serve in the attention implementation?
The causal mask ensures autoregressive generation by preventing tokens from attending to future positions in the sequence. Using torch.tril to create a lower-triangular matrix of ones, the implementation masks out upper-triangular positions (future tokens) before the softmax operation, which is essential for training language models to predict the next token based solely on previous context.
How does multi-head attention improve model performance compared to single-head attention?
Multi-head attention allows the model to jointly attend to information from different representation subspaces at different positions. By concatenating outputs from n_head parallel Head instances—each learning distinct attention patterns—the model captures richer contextual relationships than any single attention mechanism could achieve alone, followed by a learned projection matrix that mixes these representations.
What is the difference between the Head and MultiHeadAttention classes in this implementation?
The Head class implements a single attention mechanism with individual key, query, and value projections, computing scaled dot-product attention for one representation subspace. The MultiHeadAttention class orchestrates multiple Head instances (typically 4-16 heads), concatenates their outputs, and applies a final linear projection, effectively combining parallel attention mechanisms into a unified representation.
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 →