Add lesson: self-attention from scratch (Phase 7)

Query, key, value mechanism implemented with numpy. Scaled dot-product
attention, softmax, weighted value aggregation. Visualizes attention
weights on example sentences.
This commit is contained in:
Rohit Ghumare
2026-03-19 15:23:59 +05:30
parent b5c8bb0475
commit c64d63be13
3 changed files with 514 additions and 0 deletions
@@ -0,0 +1,146 @@
import numpy as np
def softmax(x):
shifted = x - np.max(x, axis=-1, keepdims=True)
exp_x = np.exp(shifted)
return exp_x / np.sum(exp_x, axis=-1, keepdims=True)
def scaled_dot_product_attention(Q, K, V):
dk = Q.shape[-1]
scores = Q @ K.T / np.sqrt(dk)
weights = softmax(scores)
output = weights @ V
return output, weights
class SelfAttention:
def __init__(self, d_model, dk, dv, seed=42):
rng = np.random.default_rng(seed)
scale_qk = np.sqrt(2.0 / (d_model + dk))
self.Wq = rng.normal(0, scale_qk, (d_model, dk))
self.Wk = rng.normal(0, scale_qk, (d_model, dk))
scale_v = np.sqrt(2.0 / (d_model + dv))
self.Wv = rng.normal(0, scale_v, (d_model, dv))
self.dk = dk
def forward(self, X):
Q = X @ self.Wq
K = X @ self.Wk
V = X @ self.Wv
return scaled_dot_product_attention(Q, K, V)
class MultiHeadSelfAttention:
def __init__(self, d_model, n_heads, seed=42):
assert d_model % n_heads == 0
self.n_heads = n_heads
self.dk = d_model // n_heads
self.dv = d_model // n_heads
self.heads = [
SelfAttention(d_model, self.dk, self.dv, seed=seed + i)
for i in range(n_heads)
]
rng = np.random.default_rng(seed + n_heads)
scale = np.sqrt(2.0 / (d_model + d_model))
self.Wo = rng.normal(0, scale, (n_heads * self.dv, d_model))
def forward(self, X):
head_outputs = []
all_weights = []
for head in self.heads:
out, w = head.forward(X)
head_outputs.append(out)
all_weights.append(w)
concatenated = np.concatenate(head_outputs, axis=-1)
output = concatenated @ self.Wo
return output, all_weights
def print_attention_matrix(weights, tokens):
print(f"\n{'':>6}", end="")
for token in tokens:
print(f"{token:>6}", end="")
print()
for i, token in enumerate(tokens):
print(f"{token:>6}", end="")
for j in range(len(tokens)):
print(f"{weights[i][j]:6.3f}", end="")
print()
def ascii_heatmap(weights, tokens, chars=" ░▒▓█"):
print(f"\n{'':>6}", end="")
for t in tokens:
print(f"{t:>6}", end="")
print()
w_max = weights.max()
for i in range(len(tokens)):
print(f"{tokens[i]:>6}", end="")
for j in range(len(tokens)):
level = int(weights[i][j] * (len(chars) - 1) / w_max)
level = min(level, len(chars) - 1)
print(f"{' ' + chars[level] + ' '}", end="")
print()
if __name__ == "__main__":
sentence = ["The", "cat", "sat", "on", "the", "mat"]
n_tokens = len(sentence)
d_model = 16
dk = 8
dv = 8
rng = np.random.default_rng(42)
X = rng.normal(0, 1, (n_tokens, d_model))
print("=" * 60)
print("SELF-ATTENTION FROM SCRATCH")
print("=" * 60)
print(f"\nSentence: {' '.join(sentence)}")
print(f"Tokens: {n_tokens}, d_model: {d_model}, dk: {dk}, dv: {dv}")
print(f"Input shape: {X.shape}")
attn = SelfAttention(d_model, dk, dv, seed=42)
output, weights = attn.forward(X)
print(f"\nOutput shape: {output.shape}")
print("\nAttention weights:")
print_attention_matrix(weights, sentence)
print("\nASCII heatmap (darker = higher attention):")
ascii_heatmap(weights, sentence)
print("\n" + "=" * 60)
print("MULTI-HEAD SELF-ATTENTION")
print("=" * 60)
n_heads = 2
mha = MultiHeadSelfAttention(d_model, n_heads, seed=42)
mha_output, head_weights = mha.forward(X)
print(f"\nHeads: {n_heads}")
print(f"Output shape: {mha_output.shape}")
for h, hw in enumerate(head_weights):
print(f"\nHead {h + 1} attention weights:")
print_attention_matrix(hw, sentence)
print("\n" + "=" * 60)
print("SOFTMAX DEMO")
print("=" * 60)
logits = np.array([2.0, 1.0, 0.1])
probs = softmax(logits)
print(f"\nLogits: {logits}")
print(f"Softmax: {probs.round(4)}")
print(f"Sum: {probs.sum():.4f}")
large_logits = np.array([100.0, 200.0, 300.0])
probs_large = softmax(large_logits)
print(f"\nLarge logits: {large_logits}")
print(f"Softmax: {probs_large.round(4)}")
print(f"Sum: {probs_large.sum():.4f}")
print("(Numerically stable - no overflow)")
@@ -0,0 +1,324 @@
# Self-Attention from Scratch
> Attention is a lookup table where every word asks "who matters to me?" - and learns the answer.
**Type:** Build
**Languages:** Python
**Prerequisites:** Phase 3 (Deep Learning Core), Phase 5 Lesson 10 (Sequence-to-Sequence)
**Time:** ~90 minutes
## The Problem
RNNs process sequences one token at a time. By the time you reach token 50, the information from token 1 has been squeezed through 50 compression steps. Long-range dependencies get crushed into a fixed-size hidden state - a bottleneck that no amount of LSTM gating fully solves.
The 2014 Bahdanau attention paper showed the fix: let the decoder look back at every encoder position and decide which ones matter for the current step. But it was still bolted onto an RNN. The 2017 "Attention Is All You Need" paper asked a sharper question: what if attention is the *only* mechanism? No recurrence. No convolution. Just attention.
Self-attention lets every position in a sequence attend to every other position in a single parallel step. That is what makes transformers fast, scalable, and dominant.
## The Concept
### The Database Lookup Analogy
Think of attention as a soft database lookup:
```
Traditional database:
Query: "capital of France" --> exact match --> "Paris"
Attention:
Query: "capital of France" --> similarity to ALL keys --> weighted blend of ALL values
```
Every token generates three vectors:
- **Query (Q)**: "What am I looking for?"
- **Key (K)**: "What do I contain?"
- **Value (V)**: "What information do I provide if selected?"
The dot product between a query and all keys produces attention scores. High score means "this key matches my query." Those scores weight the values. The output is a weighted sum of values.
### Q, K, V Computation
Each token embedding gets projected through three learned weight matrices:
```
Input embeddings (sequence of n tokens, each d-dimensional):
X = [x1, x2, x3, ..., xn] shape: (n, d)
Three weight matrices:
Wq shape: (d, dk)
Wk shape: (d, dk)
Wv shape: (d, dv)
Projections:
Q = X @ Wq shape: (n, dk) each token's query
K = X @ Wk shape: (n, dk) each token's key
V = X @ Wv shape: (n, dv) each token's value
```
Visually, for one token:
```
Wq
x_i ------[*]------> q_i "What am I looking for?"
|
| Wk
+----[*]------> k_i "What do I contain?"
|
| Wv
+----[*]------> v_i "What do I offer?"
```
### The Attention Matrix
Once you have Q, K, V for all tokens, attention scores form a matrix:
```
Scores = Q @ K^T shape: (n, n)
k1 k2 k3 k4 k5
+-----+-----+-----+-----+-----+
q1 | 2.1 | 0.3 | 0.1 | 0.8 | 0.2 | <- how much q1 attends to each key
+-----+-----+-----+-----+-----+
q2 | 0.4 | 1.9 | 0.7 | 0.1 | 0.3 |
+-----+-----+-----+-----+-----+
q3 | 0.2 | 0.6 | 2.3 | 0.5 | 0.1 |
+-----+-----+-----+-----+-----+
q4 | 0.9 | 0.1 | 0.4 | 1.7 | 0.6 |
+-----+-----+-----+-----+-----+
q5 | 0.1 | 0.3 | 0.2 | 0.5 | 2.0 |
+-----+-----+-----+-----+-----+
Each row: one token's attention over the entire sequence
```
### Why Scale?
The dot products grow with dimension dk. If dk = 64, dot products can be in the range of tens, pushing softmax into regions where gradients vanish. The fix: divide by sqrt(dk).
```
Scaled scores = (Q @ K^T) / sqrt(dk)
```
This keeps values in a range where softmax produces useful gradients.
### Softmax Turns Scores into Weights
Softmax converts raw scores into a probability distribution across each row:
```
Raw scores for q1: [2.1, 0.3, 0.1, 0.8, 0.2]
|
softmax
|
Attention weights: [0.52, 0.09, 0.07, 0.14, 0.08] (sums to ~1.0)
```
Now each token has a set of weights saying how much to attend to every other token.
### Weighted Sum of Values
The final output for each token is a weighted sum of all value vectors:
```
output_i = sum( attention_weight[i][j] * v_j for all j )
For token 1:
output_1 = 0.52 * v1 + 0.09 * v2 + 0.07 * v3 + 0.14 * v4 + 0.08 * v5
```
### Full Pipeline
```
+-------+
X (input) ----->| @ Wq |-----> Q
+-------+
+-------+
X (input) ----->| @ Wk |-----> K
+-------+ +----------+
+-------+ | |
X (input) ----->| @ Wv |-----> V ---------->| weighted |----> output
+-------+ ^ | sum |
| +----------+
+--------+--------+
| softmax |
+---------+-------+
^
+---------+-------+
| Q @ K^T / sqrt |
+-----------------+
```
Formula in one line:
```
Attention(Q, K, V) = softmax( Q @ K^T / sqrt(dk) ) @ V
```
## Build It
### Step 1: Softmax from scratch
Softmax converts raw logits into probabilities. Subtract the max for numerical stability.
```python
import numpy as np
def softmax(x):
shifted = x - np.max(x, axis=-1, keepdims=True)
exp_x = np.exp(shifted)
return exp_x / np.sum(exp_x, axis=-1, keepdims=True)
logits = np.array([2.0, 1.0, 0.1])
print(f"logits: {logits}")
print(f"softmax: {softmax(logits)}")
print(f"sum: {softmax(logits).sum():.4f}")
```
### Step 2: Scaled dot-product attention
The core function. Takes Q, K, V matrices and returns the attention output plus the weight matrix.
```python
def scaled_dot_product_attention(Q, K, V):
dk = Q.shape[-1]
scores = Q @ K.T / np.sqrt(dk)
weights = softmax(scores)
output = weights @ V
return output, weights
```
### Step 3: Self-attention class with learned projections
A full self-attention module with Wq, Wk, Wv weight matrices initialized with Xavier-like scaling.
```python
class SelfAttention:
def __init__(self, d_model, dk, dv, seed=42):
rng = np.random.default_rng(seed)
scale = np.sqrt(2.0 / (d_model + dk))
self.Wq = rng.normal(0, scale, (d_model, dk))
self.Wk = rng.normal(0, scale, (d_model, dk))
scale_v = np.sqrt(2.0 / (d_model + dv))
self.Wv = rng.normal(0, scale_v, (d_model, dv))
self.dk = dk
def forward(self, X):
Q = X @ self.Wq
K = X @ self.Wk
V = X @ self.Wv
output, weights = scaled_dot_product_attention(Q, K, V)
return output, weights
```
### Step 4: Run it on a sentence
Create fake embeddings for a sentence and watch the attention weights.
```python
sentence = ["The", "cat", "sat", "on", "the", "mat"]
n_tokens = len(sentence)
d_model = 8
dk = 4
dv = 4
rng = np.random.default_rng(42)
X = rng.normal(0, 1, (n_tokens, d_model))
attn = SelfAttention(d_model, dk, dv, seed=42)
output, weights = attn.forward(X)
print("Attention weights (each row: where that token looks):\n")
print(f"{'':>6}", end="")
for token in sentence:
print(f"{token:>6}", end="")
print()
for i, token in enumerate(sentence):
print(f"{token:>6}", end="")
for j in range(n_tokens):
w = weights[i][j]
print(f"{w:6.3f}", end="")
print()
```
### Step 5: Visualize attention with ASCII heatmap
Map attention weights to characters for a quick visual.
```python
def ascii_heatmap(weights, tokens, chars=" ░▒▓█"):
n = len(tokens)
print(f"\n{'':>6}", end="")
for t in tokens:
print(f"{t:>6}", end="")
print()
for i in range(n):
print(f"{tokens[i]:>6}", end="")
for j in range(n):
level = int(weights[i][j] * (len(chars) - 1) / weights.max())
level = min(level, len(chars) - 1)
print(f"{' ' + chars[level] + ' '}", end="")
print()
ascii_heatmap(weights, sentence)
```
## Use It
PyTorch's `nn.MultiheadAttention` does exactly what we built, plus multi-head splitting and output projection:
```python
import torch
import torch.nn as nn
d_model = 8
n_heads = 2
seq_len = 6
mha = nn.MultiheadAttention(embed_dim=d_model, num_heads=n_heads, batch_first=True)
X_torch = torch.randn(1, seq_len, d_model)
output, attn_weights = mha(X_torch, X_torch, X_torch)
print(f"Input shape: {X_torch.shape}")
print(f"Output shape: {output.shape}")
print(f"Attention weight shape: {attn_weights.shape}")
print(f"\nAttn weights (averaged over heads):")
print(attn_weights[0].detach().numpy().round(3))
```
The key difference: multi-head attention runs multiple attention functions in parallel, each with its own Q, K, V projections of size dk = d_model / n_heads, then concatenates results. This lets the model attend to different relationship types simultaneously.
## Ship It
This lesson produces:
- `outputs/prompt-attention-explainer.md` - a prompt for explaining attention through the database lookup analogy
## Exercises
1. Modify `scaled_dot_product_attention` to accept an optional mask matrix that sets certain positions to negative infinity before softmax (this is how causal/decoder masking works)
2. Implement multi-head attention from scratch: split Q, K, V into `n_heads` chunks, run attention on each, concatenate, and project through a final weight matrix Wo
3. Take two different sentences of the same length, feed them through the same SelfAttention instance, and compare their attention patterns. What changes? What stays the same?
## Key Terms
| Term | What people say | What it actually means |
|------|----------------|----------------------|
| Query (Q) | "The question vector" | A learned projection of the input that represents what information this token is looking for |
| Key (K) | "The label vector" | A learned projection that represents what information this token contains, matched against queries |
| Value (V) | "The content vector" | A learned projection carrying the actual information that gets aggregated based on attention scores |
| Scaled dot-product attention | "The attention formula" | softmax(QK^T / sqrt(dk)) @ V - scaling prevents softmax saturation in high dimensions |
| Self-attention | "The token looks at itself and others" | Attention where Q, K, V all come from the same sequence, letting every position attend to every other position |
| Attention weights | "How much focus" | A probability distribution over positions, produced by softmax over scaled dot products |
| Multi-head attention | "Parallel attention" | Running multiple attention functions with different projections, then concatenating results for richer representations |
## Further Reading
- [Attention Is All You Need (Vaswani et al., 2017)](https://arxiv.org/abs/1706.03762) - the original transformer paper
- [The Illustrated Transformer (Jay Alammar)](https://jalammar.github.io/illustrated-transformer/) - best visual walkthrough of the full architecture
- [The Annotated Transformer (Harvard NLP)](https://nlp.seas.harvard.edu/annotated-transformer/) - line-by-line PyTorch implementation with explanations
@@ -0,0 +1,44 @@
---
name: prompt-attention-explainer
description: Explain the attention mechanism through the database lookup analogy
phase: 7
lesson: 2
---
You are an expert at explaining the transformer attention mechanism. Your core teaching tool is the "database lookup" analogy.
Framework for explaining attention:
1. Start with traditional databases: a query matches a key exactly and returns one value.
2. Reframe attention as a soft database lookup:
- Query (Q): what the current token is searching for
- Key (K): what each token advertises about itself
- Value (V): the actual content each token carries
- Instead of exact match, compute similarity (dot product) between the query and ALL keys
- Instead of returning one result, return a weighted blend of ALL values
3. Walk through the math step by step:
- Q, K, V are learned linear projections of the input: Q = X @ Wq, K = X @ Wk, V = X @ Wv
- Raw scores: Q @ K^T (dot product between every query-key pair)
- Scaling: divide by sqrt(dk) to prevent softmax saturation
- Softmax: convert raw scores to a probability distribution per row
- Output: weighted sum of values using those probabilities
4. Use concrete examples. Given a sentence like "The cat sat on the mat":
- Show which tokens attend to which
- Explain why "sat" might attend strongly to "cat" (subject-verb relationship)
- Show the attention weight matrix as a grid
5. Connect to the bigger picture:
- Self-attention: Q, K, V all come from the same sequence
- Cross-attention: Q comes from one sequence, K and V from another (used in translation)
- Multi-head: multiple attention functions in parallel, each learning different relationship types
- Causal masking: preventing tokens from attending to future positions (used in GPT-style models)
Rules:
- Always show the formula: Attention(Q, K, V) = softmax(Q @ K^T / sqrt(dk)) @ V
- Use ASCII diagrams for the attention matrix when possible
- Ground every abstraction in a concrete token-level example
- Explain scaling intuitively: high-dimensional dot products produce large numbers that make softmax too peaked
- When asked about multi-head attention, explain it as "different heads learn different types of relationships: one head for syntax, another for coreference, another for positional patterns"