Implement Self-Attention from Scratch
Building Scaled Dot Product Self Attention from scratch using Query, Key, and Value linear projections in PyTorch.
Implement Self-Attention from Scratch builds the core mechanism of Transformer models. Scaled Dot-Product Attention transforms input sequence vectors into Query, Key, and Value matrices. The algorithm calculates query-key similarity dot products, scales by the square root of key dimension, applies Softmax normalization, and computes weighted sums over Value vectors.
Math Foundations
Given input sequence , Scaled Dot-Product Attention projects into three distinct matrices using learnable weight projections :
The attention mechanism computes contextual representation matrix :
- : Pairwise similarity scores between every Query token and Key token ( matrix).
- Scaling: Prevents dot products from exploding in high dimensions, preserving non-zero Softmax gradients.
- Mask: Optional Causal Mask setting future token positions to so .
Query (Q) ──┐
├──► Matrix Dot Product (Q @ K.T) ──► Scale by 1/sqrt(d_k) ──► Apply Mask ──► Softmax ──► Multiply V ──► Output
Key (K) ──┘ ▲
│
Value (V) ──────────────────────────────────────────────────────────────────────────────────────────────────┘
PyTorch Implementation from Scratch
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class SingleHeadSelfAttention(nn.Module):
def __init__(self, d_in, d_k):
super().__init__()
self.d_k = d_k
self.w_q = nn.Linear(d_in, d_k, bias=False)
self.w_k = nn.Linear(d_in, d_k, bias=False)
self.w_v = nn.Linear(d_in, d_k, bias=False)
def forward(self, x, mask=None):
# x shape: [batch_size, seq_len, d_in]
B, T, _ = x.shape
# 1. Project inputs into Query, Key, and Value matrices
q = self.w_q(x) # [B, T, d_k]
k = self.w_k(x) # [B, T, d_k]
v = self.w_v(x) # [B, T, d_k]
# 2. Compute Raw Attention Scores: Q @ K^T / sqrt(d_k)
scores = torch.bmm(q, k.transpose(1, 2)) / math.sqrt(self.d_k) # [B, T, T]
# 3. Apply Causal Mask if provided (set upper triangle to -inf)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
# 4. Softmax Normalization over Key sequence dimension
attn_weights = F.softmax(scores, dim=-1) # [B, T, T]
# 5. Weighted Average over Value vectors
output = torch.bmm(attn_weights, v) # [B, T, d_k]
return output, attn_weights
Say this out loud
Self-attention projects input tokens into Query, Key, and Value matrices. Matrix multiplication of Query and Key transpose measures token pair similarities. Scaling by the square root of key dimension prevents Softmax gradient saturation. Applying Softmax and multiplying by Value vectors produces context aware token representations.
Followups to expect
- Why does self-attention have quadratic O(T^2) time complexity? Computing the full similarity matrix between all queries and keys requires operations, creating memory bottlenecks on long text sequences.
- What is Cross Attention vs Self Attention? In Self Attention, Query, Key, and Value come from the same input sequence. In Cross Attention (Decoder), Query comes from decoder tokens while Key and Value come from encoder output tokens.
Check yourself
Question 1 of 3
What mathematical formula defines Scaled Dot-Product Attention?