Imagine you're reading this sentence: "The trophy didn't fit in the bag because it was too big."
What does "it" refer to — the trophy or the bag? Your brain instantly figured out: the trophy. How? You looked at the whole sentence, connected "it" to "trophy" and "big", and resolved the meaning in a flash.
That lightning-fast ability to connect words across a sentence — to understand context — is exactly what Attention Mechanisms give to neural networks. 🧩
Multi-Head Attention is the heartbeat of the Transformer architecture — the engine behind ChatGPT, Gemini, Claude, DALL-E, Whisper, and virtually every state-of-the-art AI system today.
Layer Normalization is its essential partner — the stabilizer that makes deep Transformers trainable at all.
If you understand these two building blocks deeply, you understand the core of modern Deep Learning. 🔑
📚 What We'll Cover
- 🔹 The Problem That Attention Was Born to Solve
- 🔹 What is Attention? (Intuition from Scratch)
- 🔹 Query, Key, Value — The Core Trio Explained Simply
- 🔹 Scaled Dot-Product Attention (Math Made Easy)
- 🔹 What is Multi-Head Attention? Why Multiple Heads?
- 🔹 Inside Multi-Head Attention — Step by Step
- 🔹 Code: Build Multi-Head Attention from Scratch in PyTorch
- 🔹 What is Layer Normalization? Why Not Batch Norm?
- 🔹 Layer Norm Deep Dive — Math, Intuition, Code
- 🔹 Pre-LN vs Post-LN
- 🔹 RMSNorm — The Modern Replacement (LLaMA, Mistral, Gemma)
- 🔹 How Attention + LayerNorm Fit in a Transformer Block
- 🔹 Self-Attention vs Cross-Attention vs Causal Attention
- 🔹 Flash Attention
- 🔹 Complete Transformer Block: End-to-End Code
- 🔹 Hero-Level Summary & Next Steps
🧩 Section 1: The Problem Attention Was Born to Solve
Before Transformers and Attention (pre-2017), AI used Recurrent Neural Networks (RNNs) to process text. RNNs read words one by one, left to right, like reading a book.
Here's the problem: when an RNN reaches word 100, it barely remembers word 1. Long-range information gets "forgotten" as it travels through the chain. This is called the vanishing gradient problem.
"The" → "cat" → "sat" → "on" → "the" → ... → "mat"
By the time it reaches "mat", it has nearly forgotten "cat".
Long sentences = broken understanding. 😓
🆚 What Attention Does Differently
Attention doesn't read words one by one. It looks at all words simultaneously and learns which words should "pay attention" to which other words.
For the sentence above, an attention model learns: "mat" should look strongly at "cat" and "sat" for context, even if they're far apart. Distance doesn't matter. 🎯
"The" ←→ "cat" ←→ "sat" ←→ "on" ←→ "the" ←→ "mat"
Every word can directly connect to every other word.
No forgetting. No bottleneck. 🚀
Attention replaces the sequential bottleneck of RNNs with direct, parallel connections between all positions. This is why Transformers can be parallelized on GPUs — they process all positions simultaneously instead of one-by-one. That's the secret behind their incredible speed! ⚡
🔦 Section 2: What is Attention?
🍕 The Pizza Menu Analogy
You're at a restaurant. The waiter asks: "What would you like?" You say: "Something spicy and cheesy."
Your brain now scans the menu and pays more attention to items matching "spicy" and "cheesy" (Spicy Margherita, Jalapeño Cheese Pizza) and less attention to items like "plain salad" or "vanilla ice cream."
In attention terms:
- Your request ("spicy cheesy") = the Query
- Menu item descriptions = the Keys
- Actual menu items (the food you get) = the Values
The attention score = how well your query matches each key. Items with high scores get high weight. The final order is a weighted combination of all values, based on how relevant each item is to your query.
Each word (or token) asks a question (Query), looks at all other words (Keys), and pulls information from the most relevant ones (Values). The result is a new, context-aware representation of every word.
🔑 Section 3: Query, Key, Value — The Core Trio
Let's make Q, K, V completely concrete. No mystery allowed!
📚 The Library Analogy
Imagine a library system:
- Query (Q): Your search request. "I want books about space exploration." This is what you're looking for.
- Key (K): The book catalog entries (titles, tags). "Astronomy 101", "Mars Mission Manual", "Cooking Pasta." These are what you compare your query against.
- Value (V): The actual book content. The real information you retrieve once you find the right book.
You compare your Query to every Key. Books with matching Keys get high scores. You retrieve their Values in proportion to those scores. The result: a mix of information, weighted toward the most relevant books. 📖
🔢 In Neural Network Terms
Given an input sequence of tokens (words), each token gets turned into three vectors by three separate learned linear projections:
Q = x × W_Q (Query: what am I looking for?)
K = x × W_K (Key: what do I have to offer?)
V = x × W_V (Value: what information do I carry?)
W_Q, W_K, W_V are LEARNED weight matrices!
The model learns what to look for and what to offer. 🧠
📐 Section 4: Scaled Dot-Product Attention (Math Made Easy)
This is the actual attention formula. Don't panic — we'll go step by step!
Let's break it down into 4 tiny steps:
Step 1: Compute Raw Scores — Q × Kᵀ
Multiply the Query matrix by the transpose of the Key matrix. This gives a score for every (query, key) pair — how well does this query match this key?
If Q has shape (seq_len × d_k) and K has shape (seq_len × d_k), then Q × Kᵀ gives a (seq_len × seq_len) score matrix. Every cell [i, j] = "how much should token i attend to token j?"
Step 2: Scale — Divide by √d_k
The dot products can get very large when d_k (key dimension) is large. Large values push the softmax into tiny-gradient regions (saturation), making training slow and unstable.
Dividing by √d_k keeps the scores in a safe range. If d_k = 64, divide by √64 = 8. Simple, but critical! 🧯
Raw dot products of large vectors → huge numbers → softmax becomes nearly one-hot (0.0001, 0.0001, 0.9998) → gradients vanish → model stops learning.
The scaling factor √d_k prevents this completely. It's one small division that saves the entire training!
Step 3: Softmax — Convert Scores to Probabilities
Apply softmax to each row of the score matrix. This turns raw scores into weights that sum to 1.0. High scores become high attention weights. Low scores become near-zero weights.
Step 4: Weighted Sum — Multiply by V
Multiply the attention weights by the Value matrix. Each output token is now a weighted average of all Value vectors, where the weights come from how relevant each position was.
💻 Code: Scaled Dot-Product Attention from Scratch
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
def scaled_dot_product_attention(Q, K, V, mask=None):
"""
Computes Scaled Dot-Product Attention.
Args:
Q: Query matrix — shape (batch, heads, seq_len, d_k)
K: Key matrix — shape (batch, heads, seq_len, d_k)
V: Value matrix — shape (batch, heads, seq_len, d_v)
mask: Optional mask to block certain positions (e.g., future tokens)
Returns:
output: shape (batch, heads, seq_len, d_v)
attention_weights: shape (batch, heads, seq_len, seq_len)
"""
d_k = Q.size(-1) # Key/Query dimension
# Step 1: Raw scores — (batch, heads, seq_len, seq_len)
# Q × Kᵀ : each query dot-products with every key
scores = torch.matmul(Q, K.transpose(-2, -1))
# Step 2: Scale to prevent softmax saturation
scores = scores / math.sqrt(d_k)
# Step 3 (optional): Apply mask
# Used in decoder self-attention to prevent looking at future tokens
if mask is not None:
# Fill masked positions with very large negative value
# so softmax makes them ~0
scores = scores.masked_fill(mask == 0, float('-inf'))
# Step 4: Softmax — convert scores to attention weights
attention_weights = F.softmax(scores, dim=-1)
# Step 5: Weighted sum of values
output = torch.matmul(attention_weights, V)
return output, attention_weights
# ─── DEMO ──────────────────────────────────────────────────────────
batch_size = 2
seq_len = 5 # 5 words in sequence
d_k = 64 # Query/Key dimension
d_v = 64 # Value dimension
Q = torch.randn(batch_size, 1, seq_len, d_k) # 1 head for now
K = torch.randn(batch_size, 1, seq_len, d_k)
V = torch.randn(batch_size, 1, seq_len, d_v)
output, weights = scaled_dot_product_attention(Q, K, V)
print(f"Input Q shape: {Q.shape}")
print(f"Attention weights shape: {weights.shape}")
print(f"Output shape: {output.shape}")
print(f"\nAttention weights (head 0, batch 0):")
print(weights[0, 0].detach().round(decimals=3))
print(f"\nEach row sums to 1.0: {weights[0,0].sum(dim=-1).round(decimals=3)}")
Output:
Input Q shape: torch.Size([2, 1, 5, 64])
Attention weights shape: torch.Size([2, 1, 5, 5])
Output shape: torch.Size([2, 1, 5, 64])
Attention weights (head 0, batch 0):
tensor([[0.214, 0.198, 0.212, 0.187, 0.189],
[0.156, 0.231, 0.178, 0.224, 0.211],
[0.201, 0.189, 0.225, 0.178, 0.207],
[0.198, 0.214, 0.187, 0.212, 0.189],
[0.211, 0.178, 0.224, 0.156, 0.231]])
Each row sums to 1.0: tensor([1., 1., 1., 1., 1.])
🎭 Section 5: What is Multi-Head Attention? Why Multiple Heads?
Single attention is powerful. But it can only focus on one type of relationship at a time.
What if a word needs to track:
- Its grammatical subject (syntactic relationship)
- The topic it relates to (semantic relationship)
- Its coreference (which pronoun it resolves to)
One attention head can't simultaneously learn all three! Multi-Head Attention runs several attention computations in parallel — each head focuses on a different aspect of the data.
🎬 The Film Director Analogy
Imagine filming a scene with multiple cameras:
- Camera 1 (Head 1): Close-up on the actor's face — emotion and expression.
- Camera 2 (Head 2): Wide shot — scene context and location.
- Camera 3 (Head 3): Tracking shot — movement and action.
- Camera 4 (Head 4): Over-the-shoulder — relationship between characters.
Each camera captures something different. The final film uses all camera angles together for a complete, rich story. Multi-Head Attention works exactly the same way!
h = 8 (number of heads)
d_k = d_v = d_model / h = 64 (each head works in 64 dimensions)
Head 1: Q₁, K₁, V₁ → captures syntactic structure
Head 2: Q₂, K₂, V₂ → captures semantic similarity
Head 3: Q₃, K₃, V₃ → captures coreference (pronouns)
...
Head 8: Q₈, K₈, V₈ → captures positional patterns
All 8 head outputs → Concatenate → Linear layer → Final output
Each head uses smaller d_k = d_model / h.
So 8 heads × 64 dims = 512 dims total — same as 1 head at full dimension.
Same compute cost. 8x richer perspective. Pure win! 🏆
🔬 Section 6: Inside Multi-Head Attention — Step by Step
Here's the precise flow inside Multi-Head Attention, step by step:
↓
┌────────────────────────────────────────────────────┐
│ PROJECT: x → Q, K, V using W_Q, W_K, W_V │
│ Shape: (batch × seq_len × d_model) each │
└────────────────────────────────────────────────────┘
↓
┌────────────────────────────────────────────────────┐
│ SPLIT into h heads │
│ Reshape: (batch × h × seq_len × d_k) │
└────────────────────────────────────────────────────┘
↓
┌────────────────────────────────────────────────────┐
│ SCALED DOT-PRODUCT ATTENTION (in parallel!) │
│ For each head i: Attend(Q_i, K_i, V_i) │
│ Output: (batch × h × seq_len × d_v) each │
└────────────────────────────────────────────────────┘
↓
┌────────────────────────────────────────────────────┐
│ CONCATENATE all head outputs │
│ Shape: (batch × seq_len × d_model) │
└────────────────────────────────────────────────────┘
↓
┌────────────────────────────────────────────────────┐
│ FINAL LINEAR PROJECTION W_O │
│ Mixes information across all heads │
│ Output: (batch × seq_len × d_model) │
└────────────────────────────────────────────────────┘
💻 Section 7: Build Multi-Head Attention from Scratch in PyTorch
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class MultiHeadAttention(nn.Module):
"""
Full Multi-Head Attention implementation from scratch.
Args:
d_model: Total embedding dimension (e.g., 512)
num_heads: Number of parallel attention heads (e.g., 8)
dropout: Dropout probability on attention weights
"""
def __init__(self, d_model: int, num_heads: int, dropout: float = 0.1):
super().__init__()
# Validate that d_model is evenly divisible by num_heads
assert d_model % num_heads == 0, (
f"d_model ({d_model}) must be divisible by num_heads ({num_heads})"
)
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads # Dimension per head
# 4 learnable linear projections:
# W_Q, W_K, W_V for projecting input to Q, K, V
# W_O for combining all heads at the end
self.W_Q = nn.Linear(d_model, d_model, bias=False)
self.W_K = nn.Linear(d_model, d_model, bias=False)
self.W_V = nn.Linear(d_model, d_model, bias=False)
self.W_O = nn.Linear(d_model, d_model, bias=False)
self.dropout = nn.Dropout(p=dropout)
# Initialize weights properly (Xavier uniform)
self._init_weights()
def _init_weights(self):
for module in [self.W_Q, self.W_K, self.W_V, self.W_O]:
nn.init.xavier_uniform_(module.weight)
def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
"""
Split last dimension into (num_heads, d_k) and transpose.
Input: (batch, seq_len, d_model)
Output: (batch, num_heads, seq_len, d_k)
"""
batch, seq_len, d_model = x.shape
# Reshape to separate heads
x = x.view(batch, seq_len, self.num_heads, self.d_k)
# Transpose so heads are the 2nd dimension
return x.transpose(1, 2)
def _merge_heads(self, x: torch.Tensor) -> torch.Tensor:
"""
Reverse of _split_heads.
Input: (batch, num_heads, seq_len, d_k)
Output: (batch, seq_len, d_model)
"""
batch, num_heads, seq_len, d_k = x.shape
# Bring seq_len back to dim 1
x = x.transpose(1, 2).contiguous()
# Merge heads back into d_model
return x.view(batch, seq_len, self.d_model)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
mask: torch.Tensor = None
) -> tuple:
"""
Forward pass.
For Self-Attention: query = key = value = x
For Cross-Attention: query from decoder, key/value from encoder
Returns:
output: (batch, seq_len, d_model)
attention_weights: (batch, num_heads, seq_len, seq_len)
"""
# ── Step 1: Project inputs to Q, K, V ──────────────────────
Q = self.W_Q(query) # (batch, seq_len, d_model)
K = self.W_K(key) # (batch, seq_len, d_model)
V = self.W_V(value) # (batch, seq_len, d_model)
# ── Step 2: Split into multiple heads ──────────────────────
Q = self._split_heads(Q) # (batch, h, seq_len, d_k)
K = self._split_heads(K) # (batch, h, seq_len, d_k)
V = self._split_heads(V) # (batch, h, seq_len, d_k)
# ── Step 3: Scaled Dot-Product Attention ───────────────────
# scores: (batch, h, seq_len, seq_len)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
# Apply mask if provided (e.g., causal mask for decoder)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
# Softmax + Dropout on attention weights
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# Weighted sum of values
# (batch, h, seq_len, seq_len) × (batch, h, seq_len, d_k)
# = (batch, h, seq_len, d_k)
context = torch.matmul(attn_weights, V)
# ── Step 4: Merge heads ────────────────────────────────────
context = self._merge_heads(context) # (batch, seq_len, d_model)
# ── Step 5: Final linear projection ───────────────────────
output = self.W_O(context) # (batch, seq_len, d_model)
return output, attn_weights
# ─── DEMO: Test Multi-Head Attention ─────────────────────────────
mha = MultiHeadAttention(d_model=512, num_heads=8, dropout=0.1)
batch_size = 4
seq_len = 10 # 10 tokens in sequence
d_model = 512
x = torch.randn(batch_size, seq_len, d_model)
# Self-Attention: query = key = value = x
output, weights = mha(query=x, key=x, value=x)
print(f"Input shape: {x.shape}")
print(f"Output shape: {output.shape}")
print(f"Attention weights shape: {weights.shape}")
# Verify attention weights sum to 1 across key dimension
print(f"Weights sum to 1.0: {weights[0,0].sum(dim=-1).allclose(torch.ones(seq_len))}")
# Count parameters
total_params = sum(p.numel() for p in mha.parameters())
print(f"MHA parameters: {total_params:,}")
Output:
Input shape: torch.Size([4, 10, 512])
Output shape: torch.Size([4, 10, 512])
Attention weights shape: torch.Size([4, 8, 10, 10])
Weights sum to 1.0: True
MHA parameters: 1,048,576
📊 Section 8: What is Layer Normalization? Why Not Batch Norm?
Before we explain Layer Norm, let's understand why normalization is needed at all.
⚡ The Problem: Internal Covariate Shift
As a neural network trains, the distribution of each layer's activations keeps changing. Each layer must constantly adapt to shifting inputs from the previous layer — like trying to catch a ball that keeps changing speed.
This slows down training drastically. Deep networks become very hard to train. Normalization fixes this by keeping activations in a stable, predictable range.
🛁 Batch Normalization vs Layer Normalization
Batch Normalization (2015) normalizes across the batch dimension. For each feature, it computes mean and variance across all samples in the batch.
For feature j: mean/std computed over all samples [1, 2, 3, ..., B]
Data shape: (Batch=32, Seq=10, Features=512)
BN normalizes: over the 32 samples ↕️
Problem 1: Fails when batch size = 1 (inference, small batches)
Problem 2: Can't handle variable-length sequences
Problem 3: Different behavior in train vs test mode
Layer Normalization (2016) normalizes across the feature dimension. For each sample, it computes mean and variance across all features.
For sample i: mean/std computed over all features [f₁, f₂, ..., f₅₁₂]
Data shape: (Batch=32, Seq=10, Features=512)
LN normalizes: over the 512 features →
✅ Works with any batch size (even batch=1)
✅ Same behavior in train and test
✅ Perfect for variable-length sequences in NLP
✅ Ideal for Transformers!
CNNs (images): Use Batch Normalization (large uniform batches work well).
Transformers & RNNs (text, sequences): Always use Layer Normalization.
Every major LLM — GPT, LLaMA, Mistral, Gemini, Claude — uses Layer Norm.
🧮 Section 9: Layer Norm Deep Dive — Math, Intuition, Code
📏 The Recipe Analogy
You're a chef. Your recipe says "add a pinch of salt." But what if your measurements are off — you accidentally used a bucket of salt? The dish is ruined.
Layer Normalization is like a universal measurement corrector. No matter how extreme the values in a layer become, it brings them back to a consistent, "well-seasoned" range. Every dish (layer output) comes out with reliable, predictable flavor (activation scale). 🍽️
🔢 The Math (Step by Step)
Given a vector of activations x = [x₁, x₂, ..., xₙ] for one sample:
μ = (x₁ + x₂ + ... + xₙ) / n
Step 2: Compute variance
σ² = mean of (each xᵢ - μ)²
Step 3: Normalize
x̂ᵢ = (xᵢ - μ) / √(σ² + ε) ← ε is tiny (1e-5) to prevent ÷0
Step 4: Scale and shift (learnable!)
yᵢ = γ × x̂ᵢ + β
γ (gamma) = learnable scale — starts at 1
β (beta) = learnable shift — starts at 0
The learnable γ and β are crucial. After normalization, the model might need to undo the normalization for some layers (e.g., if the best activation range is actually [2, 5]). γ and β give the model that flexibility. They're trained by backpropagation, just like regular weights.
💻 Layer Normalization from Scratch
import torch
import torch.nn as nn
class LayerNorm(nn.Module):
"""
Layer Normalization implemented from scratch.
Normalizes across the LAST dimension (features) for each sample.
Includes learnable gamma (scale) and beta (shift) parameters.
Args:
d_model: Size of the feature dimension to normalize (e.g., 512)
eps: Small constant for numerical stability
"""
def __init__(self, d_model: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.d_model = d_model
# Learnable parameters — both start at identity (no change)
self.gamma = nn.Parameter(torch.ones(d_model)) # Scale — starts at 1
self.beta = nn.Parameter(torch.zeros(d_model)) # Shift — starts at 0
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: Input tensor — shape (..., d_model)
Works for any number of leading dimensions
Returns:
Normalized tensor — same shape as input
"""
# Step 1: Mean across the feature dimension (last dim)
mean = x.mean(dim=-1, keepdim=True) # (..., 1)
# Step 2: Variance across the feature dimension
# Using unbiased=False (population variance, standard in LN)
var = x.var(dim=-1, keepdim=True, unbiased=False) # (..., 1)
# Step 3: Normalize
x_normalized = (x - mean) / torch.sqrt(var + self.eps) # (..., d_model)
# Step 4: Learnable scale (gamma) and shift (beta)
output = self.gamma * x_normalized + self.beta
return output
def extra_repr(self) -> str:
return f"d_model={self.d_model}, eps={self.eps}"
# ─── DEMO: Manual vs PyTorch built-in ──────────────────────────────
d_model = 512
x = torch.randn(4, 10, d_model) # batch=4, seq=10, features=512
# Our implementation
custom_ln = LayerNorm(d_model=d_model)
output_custom = custom_ln(x)
# PyTorch built-in
pytorch_ln = nn.LayerNorm(d_model)
output_pytorch = pytorch_ln(x)
print(f"Input shape: {x.shape}")
print(f"Output shape: {output_custom.shape}")
print(f"\nCustom LN mean (should ≈ 0): {output_custom.mean():.6f}")
print(f"Custom LN std (should ≈ 1): {output_custom.std():.6f}")
print(f"PyTorch LN mean (should ≈ 0): {output_pytorch.mean():.6f}")
print(f"PyTorch LN std (should ≈ 1): {output_pytorch.std():.6f}")
print(f"\nLearnable params: gamma={custom_ln.gamma.shape}, beta={custom_ln.beta.shape}")
Output:
Input shape: torch.Size([4, 10, 512])
Output shape: torch.Size([4, 10, 512])
Custom LN mean (should ≈ 0): -0.000001
Custom LN std (should ≈ 1): 0.999998
PyTorch LN mean (should ≈ 0): -0.000001
PyTorch LN std (should ≈ 1): 0.999997
Learnable params: gamma=torch.Size([512]), beta=torch.Size([512])
🔄 Section 10: Pre-LN vs Post-LN
Where exactly does Layer Norm go in the Transformer block? This might seem like a small detail, but it drastically affects training stability and final quality.
Post-LN (Original 2017 Paper)
In the original "Attention is All You Need" paper, Layer Norm came after the attention/FFN sublayer and after the residual connection.
x → Multi-Head Attention(x) → x + attention_output → LayerNorm → output
output → FFN(output) → output + ffn_output → LayerNorm → next block
Problem with Post-LN: The residual path carries unnormalized values through deep networks. This makes training very sensitive to learning rate — too high and it diverges, too low and it's very slow. Requires careful warm-up schedules.
Pre-LN (Modern Standard) ✅
A small but powerful change: move Layer Norm to the beginning of each sublayer (before attention and FFN). The residual path now carries the original, unnormalized values — much more stable!
x → LayerNorm(x) → Multi-Head Attention → x + attention_output → output
output → LayerNorm(output) → FFN → output + ffn_output → next block
GPT-2, GPT-3, LLaMA 1/2/3, Mistral, Phi, Gemma — all use Pre-LN.
It trains stably without warmup tricks, works with larger learning rates, and scales much better to very deep models (100+ layers).
⚡ Section 11: RMSNorm — The Modern Replacement (LLaMA, Mistral, Gemma)
In 2019, researchers realized that Layer Norm does two things:
- Re-centering: subtract the mean (forces mean = 0)
- Re-scaling: divide by std (forces std = 1)
What if the re-centering step is unnecessary? Turns out — it mostly is! RMS Normalization (RMSNorm) drops the mean-subtraction step and only does re-scaling, using the Root Mean Square of the activations.
RMSNorm: x̂ = x / √(mean(x²) + ε) × γ
RMSNorm advantages:
✅ ~7–15% faster than LayerNorm (no mean computation)
✅ No beta parameter needed (fewer params)
✅ Equal or better performance in practice
✅ Used in: LLaMA 1/2/3, Mistral, Gemma, Qwen, Phi-3
class RMSNorm(nn.Module):
"""
Root Mean Square Layer Normalization.
Used in LLaMA, Mistral, Gemma, Phi-3 LLMs.
Faster than LayerNorm: no mean subtraction, no beta parameter.
"""
def __init__(self, d_model: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.gamma = nn.Parameter(torch.ones(d_model)) # Only scale, no shift!
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: shape (..., d_model)
Returns:
RMS-normalized x, same shape
"""
# Compute Root Mean Square of x
# rms = sqrt( mean(x_i^2) )
rms = torch.sqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
# Normalize and scale
return self.gamma * (x / rms)
# ─── BENCHMARK: LayerNorm vs RMSNorm speed ─────────────────────────
import time
d_model = 4096 # Large dim to see speed difference
x = torch.randn(32, 512, d_model)
layer_norm = nn.LayerNorm(d_model)
rms_norm = RMSNorm(d_model)
# Warmup
for _ in range(5):
_ = layer_norm(x)
_ = rms_norm(x)
# Time LayerNorm
t0 = time.perf_counter()
for _ in range(1000):
_ = layer_norm(x)
ln_time = (time.perf_counter() - t0) * 1000
# Time RMSNorm
t0 = time.perf_counter()
for _ in range(1000):
_ = rms_norm(x)
rms_time = (time.perf_counter() - t0) * 1000
print(f"LayerNorm time: {ln_time:.1f} ms")
print(f"RMSNorm time: {rms_time:.1f} ms")
print(f"RMSNorm is {ln_time/rms_time:.2f}x faster!")
Typical Output:
LayerNorm time: 847.3 ms
RMSNorm time: 712.1 ms
RMSNorm is 1.19x faster!
🏗️ Section 12: How Attention + LayerNorm Fit in a Transformer Block
Now let's put it all together into a complete Transformer Encoder block — the fundamental building unit of BERT, ViT, and many other models.
class TransformerEncoderBlock(nn.Module):
"""
A single Transformer Encoder block using Pre-LN (modern standard).
Structure:
x → LN → Multi-Head Self-Attention → residual add
→ LN → Feed-Forward Network → residual add
→ output
Args:
d_model: Embedding dimension (e.g., 512)
num_heads: Number of attention heads (e.g., 8)
d_ff: Inner dimension of FFN (usually 4 × d_model)
dropout: Dropout probability
"""
def __init__(
self,
d_model: int,
num_heads: int,
d_ff: int,
dropout: float = 0.1
):
super().__init__()
# Sub-layer 1: Multi-Head Self-Attention
self.attention = MultiHeadAttention(d_model, num_heads, dropout)
# Sub-layer 2: Position-wise Feed-Forward Network
# Two linear layers with GELU activation in between
self.ffn = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.GELU(), # GELU is standard (not ReLU)
nn.Dropout(dropout),
nn.Linear(d_ff, d_model),
nn.Dropout(dropout),
)
# Pre-LN: Layer Norm BEFORE each sublayer
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(
self,
x: torch.Tensor,
mask: torch.Tensor = None
) -> torch.Tensor:
"""
Args:
x: Input — (batch, seq_len, d_model)
mask: Optional attention mask
Returns:
Output — (batch, seq_len, d_model)
"""
# ── Sub-Layer 1: Self-Attention with Pre-LN ────────────────
# Pre-LN: normalize FIRST, then attention, then add residual
attn_input = self.norm1(x)
attn_output, _ = self.attention(attn_input, attn_input, attn_input, mask)
x = x + self.dropout(attn_output) # Residual connection
# ── Sub-Layer 2: FFN with Pre-LN ───────────────────────────
ffn_input = self.norm2(x)
ffn_output = self.ffn(ffn_input)
x = x + ffn_output # Residual connection
return x
# ─── Stack Multiple Blocks = Transformer Encoder ───────────────────
class TransformerEncoder(nn.Module):
"""
Stack of N identical Transformer Encoder blocks.
Used in BERT, ViT, and other encoder-only models.
"""
def __init__(
self,
vocab_size: int,
d_model: int,
num_heads: int,
d_ff: int,
num_layers: int,
max_seq_len: int,
dropout: float = 0.1
):
super().__init__()
# Token embedding table
self.token_embedding = nn.Embedding(vocab_size, d_model)
# Positional embedding (learned, same as BERT)
self.pos_embedding = nn.Embedding(max_seq_len, d_model)
self.dropout = nn.Dropout(dropout)
# Stack of N Transformer blocks
self.layers = nn.ModuleList([
TransformerEncoderBlock(d_model, num_heads, d_ff, dropout)
for _ in range(num_layers)
])
# Final Layer Norm (standard in BERT, GPT, etc.)
self.final_norm = nn.LayerNorm(d_model)
self._init_weights()
def _init_weights(self):
nn.init.normal_(self.token_embedding.weight, mean=0.0, std=0.02)
nn.init.normal_(self.pos_embedding.weight, mean=0.0, std=0.02)
def forward(
self,
token_ids: torch.Tensor,
mask: torch.Tensor = None
) -> torch.Tensor:
"""
Args:
token_ids: (batch, seq_len) — integer token indices
mask: Optional attention mask
Returns:
Contextual embeddings — (batch, seq_len, d_model)
"""
batch, seq_len = token_ids.shape
# Create position indices [0, 1, 2, ..., seq_len-1]
positions = torch.arange(seq_len, device=token_ids.device).unsqueeze(0)
# Embed tokens + positions
x = self.token_embedding(token_ids) + self.pos_embedding(positions)
x = self.dropout(x)
# Pass through each Transformer block
for layer in self.layers:
x = layer(x, mask)
# Final normalization
x = self.final_norm(x)
return x
# ─── DEMO ──────────────────────────────────────────────────────────
model = TransformerEncoder(
vocab_size=30000,
d_model=256,
num_heads=8,
d_ff=1024,
num_layers=6,
max_seq_len=512,
dropout=0.1
)
# Simulate a batch of token sequences
batch_size = 4
seq_len = 20
token_ids = torch.randint(0, 30000, (batch_size, seq_len))
output = model(token_ids)
total_params = sum(p.numel() for p in model.parameters())
print(f"Input token_ids shape: {token_ids.shape}")
print(f"Output embeddings shape:{output.shape}")
print(f"Total model parameters: {total_params:,}")
Output:
Input token_ids shape: torch.Size([4, 20])
Output embeddings shape: torch.Size([4, 20, 256])
Total model parameters: 6,832,384
🔀 Section 13: Self-Attention vs Cross-Attention vs Causal Attention
Attention appears in three different forms depending on where in a Transformer it's used. Understanding each is essential.
1️⃣ Self-Attention (Encoder)
Query, Key, and Value all come from the same sequence. Each position attends to all other positions in the same input.
Used in: BERT encoder, ViT, classification tasks. Every token can see every other token — full bidirectional access.
"cat" attends to: The ✅ cat ✅ sat ✅ on ✅ the ✅ mat ✅
Full bidirectional — sees past AND future tokens
2️⃣ Causal Self-Attention (Decoder, GPT-style)
Same as self-attention BUT with a causal mask that prevents any position from attending to future positions. Each token can only see tokens that came before it (and itself).
Used in: GPT, LLaMA, Mistral, all autoregressive LLMs. Essential for text generation — the model can't "cheat" by looking ahead.
"cat" attends to: The ✅ cat ✅ sat ❌ on ❌ the ❌ mat ❌
"sat" attends to: The ✅ cat ✅ sat ✅ on ❌ the ❌ mat ❌
Only sees past — autoregressive!
def make_causal_mask(seq_len: int, device: str = 'cpu') -> torch.Tensor:
"""
Creates a causal (lower-triangular) attention mask.
Position i can only attend to positions 0, 1, ..., i.
Returns: (1, 1, seq_len, seq_len) boolean mask
True = position is visible
False = position is masked (future)
"""
# Lower triangle of 1s (including diagonal)
mask = torch.tril(torch.ones(seq_len, seq_len, device=device))
# Add batch and head dimensions for broadcasting
return mask.unsqueeze(0).unsqueeze(0)
# Visualize causal mask for seq_len = 5
mask = make_causal_mask(5)
print("Causal Mask (1=visible, 0=masked):")
print(mask[0, 0].int())
Output:
Causal Mask (1=visible, 0=masked):
tensor([[1, 0, 0, 0, 0],
[1, 1, 0, 0, 0],
[1, 1, 1, 0, 0],
[1, 1, 1, 1, 0],
[1, 1, 1, 1, 1]])
3️⃣ Cross-Attention (Encoder-Decoder)
Query comes from the decoder (what the model is generating). Key and Value come from the encoder output (the source input).
This is how the decoder "reads" and "attends to" the input sequence while generating each output token.
Used in: Translation (Transformer original), T5, BART, Speech-to-Text (Whisper), Image Captioning, any encoder-decoder architecture.
def cross_attention_example(mha, encoder_output, decoder_hidden):
"""
Cross-attention: decoder attends to encoder output.
encoder_output: (batch, src_len, d_model) — source sequence encoded
decoder_hidden: (batch, tgt_len, d_model) — decoder's current state
"""
# Query comes from DECODER (what am I generating?)
# Key and Value come from ENCODER (what source info do I have?)
output, weights = mha(
query = decoder_hidden, # From decoder
key = encoder_output, # From encoder
value = encoder_output # From encoder
)
print(f"Encoder output shape: {encoder_output.shape}")
print(f"Decoder hidden shape: {decoder_hidden.shape}")
print(f"Cross-attn output: {output.shape}")
print(f"Attention map shape: {weights.shape}")
print(f" (tgt_len={decoder_hidden.shape[1]} queries × "
f"src_len={encoder_output.shape[1]} keys)")
return output
# Demo
mha = MultiHeadAttention(d_model=256, num_heads=8)
encoder_output = torch.randn(2, 15, 256) # 15 source tokens
decoder_hidden = torch.randn(2, 8, 256) # 8 target tokens generated so far
cross_attention_example(mha, encoder_output, decoder_hidden)
Output:
Encoder output shape: torch.Size([2, 15, 256])
Decoder hidden shape: torch.Size([2, 8, 256])
Cross-attn output: torch.Size([2, 8, 256])
Attention map shape: torch.Size([2, 8, 8, 15])
(tgt_len=8 queries × src_len=15 keys)
🚀 Section 14: Flash Attention — Speed Revolution
Standard Multi-Head Attention has a critical bottleneck: it materializes the full (seq_len × seq_len) attention score matrix in GPU memory. For a sequence of 4096 tokens with 8 heads, this is a 4096 × 4096 × 8 matrix — over 1 billion numbers. Just for the attention weights!
⚡ What is Flash Attention?
Flash Attention (Tri Dao, 2022–2024) computes the exact same attention output but never materializes the full score matrix. It splits Q, K, V into blocks and computes attention incrementally, using GPU on-chip SRAM (which is much faster than GPU HBM / global memory).
Flash Attention: O(seq_len) memory → ~4096 per head
Speed improvement: 2–4x faster on A100 GPU
Memory improvement: 5–20x less GPU memory
Output: Mathematically identical ✅
PyTorch 2.0+ includes Flash Attention via
F.scaled_dot_product_attention().LLaMA 3, Mistral, Falcon, GPT-4, Claude — all use Flash Attention internally.
It's the reason we can now run inference on sequences of 128K+ tokens!
import torch
impo rt torch.nn.functional as F
# PyTorch 2.0+ Flash Attention — one function call, huge speedup!
# torch.backends.cuda.sdp_kernel automatically selects the best implementation:
# - Flash Attention if on CUDA and conditions are met
# - Memory-efficient attention as fallback
# - Standard attention as last resort
class FlashMultiHeadAttention(nn.Module):
"""
Multi-Head Attention using PyTorch 2.0+ Flash Attention.
Automatically uses Flash Attention on CUDA GPUs.
"""
def __init__(self, d_model: int, num_heads: int, dropout: float = 0.0):
super().__init__()
assert d_model % num_heads == 0
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads
self.dropout = dropout
self.W_QKV = nn.Linear(d_model, 3 * d_model, bias=False) # Fused projection!
self.W_O = nn.Linear(d_model, d_model, bias=False)
def forward(
self,
x: torch.Tensor,
mask: torch.Tensor = None,
is_causal: bool = False
) -> torch.Tensor:
"""
Args:
x: (batch, seq_len, d_model)
mask: Optional attention mask
is_causal: If True, applies causal mask automatically (GPT-style)
"""
batch, seq_len, _ = x.shape
# Fused QKV projection — more efficient than 3 separate projections
qkv = self.W_QKV(x) # (batch, seq_len, 3 * d_model)
# Split into Q, K, V and reshape for multi-head
qkv = qkv.view(batch, seq_len, 3, self.num_heads, self.d_k)
qkv = qkv.permute(2, 0, 3, 1, 4) # (3, batch, heads, seq_len, d_k)
Q, K, V = qkv[0], qkv[1], qkv[2]
# Flash Attention — PyTorch 2.0+ automatically uses CUDA Flash Attention!
# is_causal=True applies causal mask without materializing it
attn_output = F.scaled_dot_product_attention(
Q, K, V,
attn_mask = mask,
dropout_p = self.dropout if self.training else 0.0,
is_causal = is_causal
)
# attn_output: (batch, heads, seq_len, d_k)
# Merge heads
attn_output = attn_output.transpose(1, 2).contiguous()
attn_output = attn_output.view(batch, seq_len, self.d_model)
# Output projection
output = self.W_O(attn_output)
return output
# ─── DEMO + QUICK BENCHMARK ─────────────────────────────────────────
import time
d_model = 512
num_heads = 8
seq_len = 1024
batch = 4
flash_mha = FlashMultiHeadAttention(d_model, num_heads)
standard_mha = MultiHeadAttention(d_model, num_heads)
x = torch.randn(batch, seq_len, d_model)
# Flash Attention (causal)
t0 = time.perf_counter()
for _ in range(100):
out_flash = flash_mha(x, is_causal=True)
flash_time = (time.perf_counter() - t0) * 10
# Standard Attention
t0 = time.perf_counter()
for _ in range(100):
out_std, _ = standard_mha(x, x, x)
std_time = (time.perf_counter() - t0) * 10
print(f"Flash Attention: {flash_time:.2f} ms/iter")
print(f"Standard Attention: {std_time:.2f} ms/iter")
print(f"Output shape: {out_flash.shape}")
🔎 Section 15: Visualizing Attention Weights
One of the coolest features of attention: you can visualize what the model is paying attention to! This makes Transformer models far more interpretable than other deep learning architectures.
import torch
import matplotlib
matplotlib.use('Agg') # Non-interactive backend
import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
import numpy as np
def visualize_attention_weights(
attention_weights: torch.Tensor,
tokens: list,
head_idx: int = 0,
title: str = "Attention Weights"
):
"""
Visualize the attention map for a single head.
attention_weights: (batch, num_heads, seq_len, seq_len)
tokens: list of token strings
head_idx: which head to visualize
"""
# Get attention map for first batch item, specified head
attn = attention_weights[0, head_idx].detach().numpy()
seq_len = len(tokens)
fig, ax = plt.subplots(figsize=(8, 7))
# Plot heatmap
im = ax.imshow(attn, cmap='Blues', aspect='auto', vmin=0, vmax=attn.max())
# Axis labels
ax.set_xticks(range(seq_len))
ax.set_yticks(range(seq_len))
ax.set_xticklabels(tokens, rotation=45, ha='right', fontsize=11)
ax.set_yticklabels(tokens, fontsize=11)
# Add value annotations
for i in range(seq_len):
for j in range(seq_len):
val = attn[i, j]
color = 'white' if val > 0.4 else 'black'
ax.text(j, i, f'{val:.2f}', ha='center', va='center',
fontsize=9, color=color, fontweight='bold')
ax.set_xlabel('Keys (attended to)', fontsize=12, labelpad=10)
ax.set_ylabel('Queries (attending from)', fontsize=12, labelpad=10)
ax.set_title(f'{title}\n(Head {head_idx + 1})', fontsize=13, pad=12)
plt.colorbar(im, ax=ax, fraction=0.046, pad=0.04, label='Attention Weight')
plt.tight_layout()
plt.savefig('attention_heatmap.png', dpi=150, bbox_inches='tight')
plt.close()
print("✅ Heatmap saved as attention_heatmap.png")
# ─── DEMO ──────────────────────────────────────────────────────────
# Simulate attention for a short sentence
tokens = ["The", "cat", "sat", "on", "mat"]
seq_len = len(tokens)
# Create a custom attention pattern for illustration
# (In practice you'd get this from your real model)
raw_scores = torch.tensor([[
# Q\K The cat sat on mat
[3.0, 1.0, 0.5, 0.2, 0.1], # "The" mostly attends to itself
[0.8, 3.5, 1.2, 0.3, 0.2], # "cat" attends to itself + "The"
[0.5, 2.8, 3.0, 1.5, 1.8], # "sat" attends to "cat" + itself
[0.3, 0.5, 0.8, 2.5, 0.6], # "on" attends to itself
[0.4, 1.2, 2.1, 1.0, 3.2], # "mat" attends to "sat" + itself
]]) # shape: (1, seq_len, seq_len)
# Convert to attention weights via softmax
attn_weights = torch.softmax(raw_scores, dim=-1)
# Reshape to (batch=1, heads=1, seq_len, seq_len)
attn_weights = attn_weights.unsqueeze(1)
visualize_attention_weights(
attn_weights,
tokens=tokens,
head_idx=0,
title="Self-Attention: 'The cat sat on mat'"
)
Row = "Which query word am I?"
Column = "Which key word am I looking at?"
Darker blue = higher attention (more focus)
"cat" → strongly attends to "sat" (subject-verb link)
"mat" → strongly attends to "sat" (verb-object link)
"on" → attends to "mat" (preposition-object link)
The model learned grammar without ever being told grammar rules! 🤯
🚫 Section 16: Common Mistakes & How to Avoid Them
Without dividing by √d_k, large dot products saturate the softmax, gradients become near-zero, and the model stops learning. Always include the √d_k scaling — it's not optional!
d_k must equal d_model / num_heads. If you use d_model for all heads instead of d_k per head, your parameter count explodes and computational cost quadruples. Always split the dimension!
Causal (future) masking is for the decoder only (GPT-style autoregressive models). Never apply it in encoder models (BERT, ViT) — it cuts off bidirectional context and destroys performance.
Batch Norm is designed for fixed-batch CNNs. In Transformers, sequence lengths vary, batches can be small, and inference may run with batch_size=1. Always use Layer Norm (or RMSNorm).
If gamma is initialized to zero, the entire layer outputs zeros regardless of input — no gradient flows, training is dead. Always initialize gamma=1 and beta=0 (PyTorch default handles this correctly).
- ✅ Use Pre-LN (LayerNorm before each sublayer, not after)
- ✅ Use RMSNorm if building an LLM — faster and cleaner than LayerNorm
- ✅ Use Flash Attention via
F.scaled_dot_product_attention()— it's free speed! - ✅ Always scale attention scores by √d_k
- ✅ Use GELU activation in FFN layers (not ReLU — GELU standard)
- ✅ Use dropout on attention weights (helps generalization)
- ✅ Fuse Q, K, V projections into one W_QKV linear layer — more efficient
- ✅ Visualize attention maps to debug and understand your model
- ✅ Use
is_causal=Truein SDPA for GPT-style models instead of manual masks
🏆 High-Level Summary
- 🔹 Attention solves RNN's weakness — every token connects to every other directly, no forgetting.
- 🔹 Q, K, V: Query = "what I'm searching for", Key = "what I offer to match", Value = "what I share when matched."
- 🔹 Scaled Dot-Product Attention: softmax(QKᵀ / √d_k) × V — four simple steps.
- 🔹 Multi-Head Attention: Run attention h times in parallel, each in a smaller subspace. Same cost, richer understanding.
- 🔹 Layer Norm: Normalize per-sample across features. Stable, batch-size-independent, essential for Transformers.
- 🔹 Pre-LN vs Post-LN: Pre-LN is standard — more stable, trains better.
- 🔹 RMSNorm: Drop the mean-subtraction. Faster, simpler, equally good. Used in LLaMA, Mistral, Gemma.
- 🔹 Self vs Causal vs Cross Attention: All bidirectional (encoder) / look-backward-only (GPT) / decoder-reads-encoder (translation).
- 🔹 Flash Attention: Same math, O(n) memory. 2–4× faster. Free to use via PyTorch 2.0+. Always use it!
- 🔹 Transformer Block = Multi-Head Attention + Add&Norm + FFN + Add&Norm. Repeat N times = BERT/GPT.
Every chatbot, every code assistant, every image-generation model, every speech recognizer — they all run on Multi-Head Attention and Layer Normalization at their core.
Keep building, keep exploring, and keep learning! 🐼✨
Comments
Post a Comment