USAAIO Lesson 79, from Phase 3 on transformers and NLP. It starts from the sequence bottleneck in an RNN, then covers Bahdanau alignment scores and self-attention, where Q, K, and V all come from the same sequence. It derives Attention(Q,K,V) = softmax(QK^T/sqrt(d_k))V, proves by a variance argument that dot products grow with d_k, and explains why the scaling factor prevents the softmax from saturating. It was verified with torch 2.7.1+cpu in June 2026. The lesson runs to 28 slides.
Subject: Machine Learning · 54 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 79 · Phase 3 (Transformers & NLP)
From the RNN bottleneck to Bahdanau alignment scores, self-attention, and the full derivation of Attention(Q,K,V) = softmax(QKᵀ / √d_k) V — with a proof that the scaling factor is not optional.
Objectives
Var(qᵀk) = d_k when q, k ~ N(0,1), and show this causes softmax saturationWarm-up
Discussion prompt
Before we open Lesson 79: Scaled Dot-Product Attention: without looking back, what was the main idea of Mock Exam — Phase 2 (Theory + Coding), and what could you do by the end of it that you could not do before?
Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.
Answer:
3-hour combined Phase 2 mock exam — 30 theory questions spanning all Phase 2 topics, plus two timed coding problems: logistic regression from scratch and a CNN in PyTorch, each with a theory sub-question and a verified answer key. The Phase-2 coding and theory milestone.
Section
Part 1 of 3
Concept
A standard RNN encoder-decoder (Sutskever 2014) maps an entire source sentence into one fixed-length vector — the final hidden state h_T — from which the decoder must recover everything.
\[ h_t = f(h_{t-1},\, x_t), \quad c = h_T \]
For long sentences the single vector c is asked to hold too much information. Translation quality degrades sharply past ~20 tokens (Bahdanau 2015, Fig. 2). This is the bottleneck.
Counterexample
Discussion prompt
A standard RNN encoder-decoder (Sutskever 2014) maps an entire source sentence into one fixed-length vector — the final hidden state h_T — from which the decoder must recover everything.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Answer:
For long sentences the single vector c is asked to hold too much information. Translation quality degrades sharply past ~20 tokens (Bahdanau 2015, Fig. 2). This is the bottleneck.
Concept
Bahdanau's fix: let the decoder compute a soft alignment over all encoder states at each decode step, instead of using a single c.
\[ e_{tj} = a(s_{t-1},\, h_j), \quad \alpha_{tj} = \frac{e^{e_{tj}}}{\sum_k e^{e_{tk}}}, \quad c_t = \sum_j \alpha_{tj}\, h_j \]
| symbol | meaning |
|---|---|
| s_{t-1} | decoder hidden state at previous step |
| h_j | encoder hidden state at source position j |
| e_tj | alignment score (MLP or dot product) |
| α_tj | attention weight (sums to 1 over j) |
| c_t | context vector — weighted sum of encoder states |
Comparison
Comparison matrix
From Bahdanau attention (2015): alignment scores: refill the meaning column from what you know. The rest of the table is as it appeared.
| symbol | meaning |
|---|---|
| s_{t-1} | decoder hidden state at previous step |
| h_j | encoder hidden state at source position j |
| e_tj | alignment score (MLP or dot product) |
| α_tj | attention weight (sums to 1 over j) |
| c_t | context vector — weighted sum of encoder states |
Section
Part 2 of 3
Concept
Transformer attention (Vaswani 2017) abstracts Bahdanau into three matrices: a query Q (what am I looking for?), keys K (what does each entry advertise?), and values V (what each entry contributes if selected).
\[ \text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V \]
| matrix | shape | role |
|---|---|---|
| Q | (T_q, d_k) | one row per query token — what it seeks |
| K | (T_k, d_k) | one row per key token — what it offers |
| V | (T_k, d_v) | one row per value token — what it contributes |
| output | (T_q, d_v) | weighted mixture of value rows |
Trade off
Comparison matrix
From Q, K, V: the general frame: every row here is a choice with a cost. Fill the shape column, then say which row you would actually pick and what you give up for it.
| matrix | shape | role |
|---|---|---|
| Q | (T_q, d_k) | one row per query token — what it seeks |
| K | (T_k, d_k) | one row per key token — what it offers |
| V | (T_k, d_v) | one row per value token — what it contributes |
| output | (T_q, d_v) | weighted mixture of value rows |
Concept
In self-attention all three matrices are projections of the same input sequence X. Every token simultaneously acts as query, key, and value — it asks and answers its own question about every other token.
\[ Q = XW_Q,\quad K = XW_K,\quad V = XW_V \quad (X\in\mathbb{R}^{T\times d_{\rm model}}) \]
This lets token i attend to tokens j anywhere in the sequence, regardless of distance — the key advantage over RNNs (Lesson 70 context: RNNs lose long-range dependencies through the vanishing gradient).
Estimation
Predict first
Trace the full computation for a 3-token sequence with d_k=4. Seed 42, Q and K fixed to torch.randn(3,4).
Commit before you compute: what does Traced attention forward pass (T=3, d_k=4) come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: scores[0] = [0.0725, 0.099, 0.325]; after softmax → weights[0] = [0.3017, 0.3098, 0.3884]
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Higher score for token 2 → it receives the largest weight in token 0's output mix.
Worked example
Trace the full computation for a 3-token sequence with d_k=4. Seed 42, Q and K fixed to torch.randn(3,4).
import torch, torch.nn.functional as F
import numpy as np
torch.manual_seed(42)
d_k = 4
Q = torch.randn(3, d_k)
K = torch.randn(3, d_k)
V = torch.tensor([[ 0.1, -0.5, 0.8, 0.3],
[ 0.6, 0.2, -0.4, 0.9],
[-0.3, 0.7, 0.1, -0.2]])
scores = Q @ K.T / (d_k ** 0.5)
weights = F.softmax(scores, dim=-1)
out = weights @ V
print('scores :', np.round(scores.numpy(), 4))
print('weights:', np.round(weights.numpy(), 4))
print('output :', np.round(out.numpy(), 4))scores[0] = [0.0725, 0.099, 0.325]; after softmax → weights[0] = [0.3017, 0.3098, 0.3884]
Why: Higher score for token 2 → it receives the largest weight in token 0's output mix. The weights sum to 1 (they are a probability distribution over value rows).
| token i | score row (3 values) | weight row (sums to 1) | which token dominates |
|---|---|---|---|
| 0 | [0.0725, 0.099, 0.325] | [0.302, 0.310, 0.388] | token 2 (barely) |
| 1 | [-1.863, -1.425, -1.439] | [0.245, 0.380, 0.375] | tokens 1 & 2 |
| 2 | [0.154, -0.094, 0.638] | [0.294, 0.229, 0.477] | token 2 |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
scores[0] = [0.0725, 0.099, 0.325]; after softmax → weights[0] = [0.3017, 0.3098, 0.3884]
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Trace the full computation for a 3-token sequence with d_k=4. Seed 42, Q and K fixed to torch.randn(3,4).
Estimation
Predict first
Verify the homework claim: for self-attention (Q=K=V=X), output row i equals ∑_j W[i,j] · V[j].
Commit before you compute: what does Output is a weighted average of value rows come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: W[0] = [0.499, 0.039, 0.309, 0.153]; sum = 1.000; manual == out[0]: True
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. The softmax weights are a valid probability distribution over the T=4 value rows.
Worked example
Verify the homework claim: for self-attention (Q=K=V=X), output row i equals ∑_j W[i,j] · V[j].
import torch, torch.nn.functional as F
import numpy as np
torch.manual_seed(7)
X = torch.randn(4, 8) # 4 tokens, d=8
W = F.softmax(X @ X.T / 8**0.5, dim=-1) # (4,4)
out = W @ X # (4,8)
# manual check for token 0
manual0 = sum(W[0, j].item() * X[j] for j in range(4))
print('W[0]:', np.round(W[0].numpy(), 4))
print('sum W[0]:', W[0].sum().item())
print('match:', torch.allclose(manual0, out[0], atol=1e-5))W[0] = [0.499, 0.039, 0.309, 0.153]; sum = 1.000; manual == out[0]: True
Why: The softmax weights are a valid probability distribution over the T=4 value rows. The output is exactly their weighted sum — attention interpolates between value vectors by learned similarity.
| j | W[0,j] | contribution |
|---|---|---|
| 0 | 0.4990 | dominates — token 0 attends mostly to itself |
| 1 | 0.0390 | barely contributes |
| 2 | 0.3093 | second largest weight |
| 3 | 0.1527 | minor contribution |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
W[0] = [0.499, 0.039, 0.309, 0.153]; sum = 1.000; manual == out[0]: True
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Verify the homework claim: for self-attention (Q=K=V=X), output row i equals ∑_j W[i,j] · V[j].
Anomaly
Predict first
A student writes this, and it looks reasonable:
In encoder self-attention, Q comes from the encoder and K, V come from the decoder — that's why it's called 'self' (the encoder queries itself).
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Backwards. In self-attention Q, K, V all come from the SAME sequence (encoder attends to encoder, decoder attends to decoder).
Self-attention: Q, K, V are all projections of the same X. Cross-attention (in the transformer decoder): Q from decoder states, K and V from encoder outputs.
Why: Backwards. In self-attention Q, K, V all come from the SAME sequence (encoder attends to encoder, decoder attends to decoder). Cross-attention is different: Q from decoder, K/V from encoder.
Trap
In encoder self-attention, Q comes from the encoder and K, V come from the decoder — that's why it's called 'self' (the encoder queries itself).
Use Q from encoder, K and V from decoder
Why: Backwards. In self-attention Q, K, V all come from the SAME sequence (encoder attends to encoder, decoder attends to decoder). Cross-attention is different: Q from decoder, K/V from encoder.
Self-attention: Q, K, V are all projections of the same X. Cross-attention (in the transformer decoder): Q from decoder states, K and V from encoder outputs.
Self: Q=XW_Q, K=XW_K, V=XW_V (one sequence). Cross: Q=X_dec·W_Q, K=X_enc·W_K, V=X_enc·W_V
Why: The word 'self' means Q, K, V all derive from the same source sequence. Bahdanau attention is actually cross-attention: decoder queries encoder states.
Break the constraint
Discussion prompt
The rule this trap just fixed:
The word 'self' means Q, K, V all derive from the same source sequence. Bahdanau attention is actually cross-attention: decoder queries encoder states.
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
Backwards. In self-attention Q, K, V all come from the SAME sequence (encoder attends to encoder, decoder attends to decoder). Cross-attention is different: Q from decoder, K/V from encoder.
Section
Part 3 of 3
Concept
If q and k are independent N(0,1) vectors of dimension d_k, their dot product qᵀk = ∑ q_i k_i has zero mean — but its variance grows linearly with d_k.
\[ \text{Var}(q^\top k) = \sum_{i=1}^{d_k} \text{Var}(q_i k_i) = \sum_{i=1}^{d_k} 1 = d_k \]
Large variance → large scores → softmax pushes all weight onto one entry (saturation) → gradient through softmax ≈ 0. Dividing by √d_k rescales the dot product to unit variance.
Analogy
Discussion prompt
Explain Why divide by √d_k? by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.
Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.
Answer:
If q and k are independent N(0,1) vectors of dimension d_k, their dot product qᵀk = ∑ q_i k_i has zero mean — but its variance grows linearly with d_k.
Estimation
Predict first
Measure the entropy of the softmax distribution as d_k grows, with and without √d_k scaling. High entropy = spread attention; low entropy ≈ one-hot = saturated.
Commit before you compute: what does Proving saturation numerically come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: d=256: raw_std=22.2 (H=0.000, completely saturated), scaled_std=1.39 (H=1.065, still spread)
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Without scaling, by d_k=256 the softmax has entirely collapsed to a one-hot — gradient is zero and the model can't learn.
Worked example
Measure the entropy of the softmax distribution as d_k grows, with and without √d_k scaling. High entropy = spread attention; low entropy ≈ one-hot = saturated.
import torch, torch.nn.functional as F
results = []
for d in [4, 16, 64, 256]:
torch.manual_seed(0)
q = torch.randn(1, d); k = torch.randn(5, d)
raw = (q @ k.T)[0]
scaled = raw / d**0.5
def entropy(x): p=F.softmax(x,dim=-1); return -(p*torch.log(p+1e-9)).sum().item()
results.append((d, raw.std().item(), scaled.std().item(),
entropy(raw), entropy(scaled)))
for d, rs, ss, hr, hs in results:
print(f'd={d:3d}: raw_std={rs:.3f}, scaled_std={ss:.3f}, H_raw={hr:.3f}, H_scaled={hs:.3f}')d=256: raw_std=22.2 (H=0.000, completely saturated), scaled_std=1.39 (H=1.065, still spread)
Why: Without scaling, by d_k=256 the softmax has entirely collapsed to a one-hot — gradient is zero and the model can't learn. Scaling keeps std near 1 and entropy high regardless of dimension.
| d_k | raw score std | scaled std | H(unscaled) | H(scaled) |
|---|---|---|---|---|
| 4 | 1.240 | 0.620 | 1.240 | 1.491 |
| 16 | 3.436 | 0.859 | 0.626 | 1.439 |
| 64 | 7.321 | 0.915 | 0.404 | 1.367 |
| 256 | 22.189 | 1.387 | 0.000 | 1.065 |
Pattern
Step through it
Step through Proving saturation numerically one row at a time. What is driving the change, and what would the row after the last one be?
Concept
Let q_i, k_i ~ N(0,1) i.i.d. Then q_i k_i has E[q_i k_i] = 0 and Var(q_i k_i) = E[q_i² k_i²] − 0 = 1·1 = 1 (independence + unit normal second moment).
\[ q^\top k = \sum_{i=1}^{d_k} q_i k_i \quad\Rightarrow\quad \text{Var}(q^\top k) = \sum_{i=1}^{d_k} 1 = d_k \]
\[ \frac{q^\top k}{\sqrt{d_k}} \quad\Rightarrow\quad \text{Var} = \frac{d_k}{d_k} = 1 \]
Sorting
Sort into buckets
These are the pieces of Lesson 79: Scaled Dot-Product Attention, out of order. Put each one back under the part of the lesson it belongs to.
Anomaly
Predict first
A student writes this, and it looks reasonable:
The 1/√d_k factor is just a normalization convention — it doesn't change which token gets the highest weight, so it doesn't really affect training.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Softmax is NOT argmax. Rescaling changes the sharpness of the distribution: large scores produce near-one-hot outputs.
Scaling controls the softmax temperature. Without it, at d_k=256 the entropy collapses to 0.000 — the gradient of softmax w.r.t. logits is p_i(1-p_i), near zero when p_i ≈ 1.
Why: Softmax is NOT argmax. Rescaling changes the sharpness of the distribution: large scores produce near-one-hot outputs. That makes gradients nearly zero and breaks backprop through the attention weights.
Trap
The 1/√d_k factor is just a normalization convention — it doesn't change which token gets the highest weight, so it doesn't really affect training.
Omit scaling, argue argmax is invariant to positive rescaling
Why: Softmax is NOT argmax. Rescaling changes the sharpness of the distribution: large scores produce near-one-hot outputs. That makes gradients nearly zero and breaks backprop through the attention weights.
Scaling controls the softmax temperature. Without it, at d_k=256 the entropy collapses to 0.000 — the gradient of softmax w.r.t. logits is p_i(1-p_i), near zero when p_i ≈ 1.
Run the entropy table above: d_k=256, no scaling → H=0.000; with scaling → H=1.065
Why: Gradient of softmax ∂α_j/∂e_j = α_j(1-α_j). If α_j→1, gradient→0; if α_j≈1/T, gradient is maximized. The scaling keeps the regime trainable.
Two truths and a lie
Sort into buckets
Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.
c.; If q and k are independent N(0,1) vectors of dimension d_k, their dot product qᵀk = ∑ q_i k_i has zero mean — but its variance grows linearly with d_k.; Let q_i, k_i ~ N(0,1) i.i.d. Then q_i k_i has E[q_i k_i] = 0 and Var(q_i k_i) = E[q_i² k_i²] − 0 = 1·1 = 1 (independence + unit normal second moment).1/√d_k factor is just a normalization convention — it doesn't change which token gets the highest weight, so it doesn't really affect training.Concept
A hard lookup table maps a key exactly to one value: output = V[k == query]. Attention is the soft version: it computes similarity to every key and returns a weighted mixture of all values.
This is precisely Nadaraya-Watson kernel regression: f(x) = ∑_i K(x, x_i) · y_i / ∑_i K(x, x_i) where the kernel K = exp(qᵀk/√d_k) after softmax normalization.
| view | query | keys | values | weighting |
|---|---|---|---|---|
| hard lookup | exact key | table keys | table values | one-hot |
| soft attention | Q row | K rows (all) | V rows (all) | softmax similarities |
| kernel regression | x | x_i (all) | y_i (all) | normalized kernel K(x,x_i) |
Comparison
Comparison matrix
From Attention as soft lookup — and kernel regression: refill the keys column from what you know. The rest of the table is as it appeared.
| view | query | keys | values | weighting |
|---|---|---|---|---|
| hard lookup | exact key | table keys | table values | one-hot |
| soft attention | Q row | K rows (all) | V rows (all) | softmax similarities |
| kernel regression | x | x_i (all) | y_i (all) | normalized kernel K(x,x_i) |
Ranking
Put in order
These are the steps of Scaled dot-product attention: the full recipe, scrambled. Put them back in order before the next slide shows you.
S = QKᵀ — each query's dot product with every key; shape (T_q × T_k)S / √d_k — ensures Var(scores) ≈ 1, prevents softmax saturationW = softmax(S/√d_k, dim=-1) — each row is a probability distribution over keysOutput = WV — weighted average of value rows; shape (T_q × d_v)Why: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.
Pattern
S = QKᵀ — each query's dot product with every key; shape (T_q × T_k)S / √d_k — ensures Var(scores) ≈ 1, prevents softmax saturationW = softmax(S/√d_k, dim=-1) — each row is a probability distribution over keysOutput = WV — weighted average of value rows; shape (T_q × d_v)\[ \text{Attention}(Q,K,V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V \]
Edge cases
Discussion prompt
Scaled dot-product attention: the full recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
S = QKᵀ — each query's dot product with every key; shape (T_q × T_k)S / √d_k — ensures Var(scores) ≈ 1, prevents softmax saturationW = softmax(S/√d_k, dim=-1) — each row is a probability distribution over keysOutput = WV — weighted average of value rows; shape (T_q × d_v)Elimination
Eliminate the wrong options
If q and k are independent N(0,1) vectors of dimension d_k, what is Var(qᵀk)?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: Var(qᵀk) = Var(∑ q_i k_i) = ∑ Var(q_i k_i) = ∑ 1 = d_k. Each component q_i k_i has Var=1 (independence + E[z²]=1 for z~N(0,1)), and they sum. Dividing by √d_k gives a rescaled dot product with variance d_k/d_k = 1.
Check
Work out the variance claim before clicking.
Check your understanding
If q and k are independent N(0,1) vectors of dimension d_k, what is Var(qᵀk)?
Answer: A
Why: Var(qᵀk) = Var(∑ q_i k_i) = ∑ Var(q_i k_i) = ∑ 1 = d_k. Each component q_i k_i has Var=1 (independence + E[z²]=1 for z~N(0,1)), and they sum. Dividing by √d_k gives a rescaled dot product with variance d_k/d_k = 1.
Prediction
Predict first
In self-attention with Q=K=V=X (no learned projections), the i-th output vector is:
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: a weighted average of the rows of X, with weights softmax(X[i]·Xᵀ / √d_k)
Why: Output[i] = ∑_j W[i,j] · X[j] where W[i,j] = softmax_j(X[i]·X[j] / √d_k). Verified in the deck: W[0]=[0.499,0.039,0.309,0.153], and the manual weighted sum equals out[0] exactly.
Check
Think through what the matrix multiply WV actually computes.
Check your understanding
In self-attention with Q=K=V=X (no learned projections), the i-th output vector is:
Answer: A
Why: Output[i] = ∑_j W[i,j] · X[j] where W[i,j] = softmax_j(X[i]·X[j] / √d_k). Verified in the deck: W[0]=[0.499,0.039,0.309,0.153], and the manual weighted sum equals out[0] exactly.
Elimination
Eliminate the wrong options
How does scaled dot-product attention differ from a hard lookup table?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: A lookup table is a one-hot selection (only the matching key's value is returned). Attention is differentiable and returns ∑_j softmax(score_j) · V[j] — a continuous interpolation. This is exactly Nadaraya-Watson kernel regression with an exponential kernel.
Check
Recall the kernel regression analogy.
Check your understanding
How does scaled dot-product attention differ from a hard lookup table?
Answer: A
Why: A lookup table is a one-hot selection (only the matching key's value is returned). Attention is differentiable and returns ∑_j softmax(score_j) · V[j] — a continuous interpolation. This is exactly Nadaraya-Watson kernel regression with an exponential kernel.
Section
Project
Concept
Implement attention(Q, K, V, d_k) from scratch in PyTorch, verify the output matches torch.nn.functional.scaled_dot_product_attention, and demonstrate softmax saturation with and without scaling.
| # | milestone | key tool |
|---|---|---|
| 1 | Implement attention formula; verify row sums = 1 | Q@K.T, F.softmax, @V |
| 2 | Prove Var(q^T k) = d_k numerically for d in [4,16,64] | torch.randn, .var() |
| 3 | Show entropy collapse without scaling; full program | F.softmax entropy |
Build rules: type every line, print shapes at each step, verify weight rows sum to 1.0 using assert abs(W.sum(-1) - 1).max() < 1e-5.
Counterexample
Discussion prompt
Implement attention(Q, K, V, d_k) from scratch in PyTorch, verify the output matches torch.nn.functional.scaled_dot_product_attention, and demonstrate softmax saturation with and without scaling.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Answer:
Build rules: type every line, print shapes at each step, verify weight rows sum to 1.0 using assert abs(W.sum(-1) - 1).max() < 1e-5.
Worked example
Your turn: write attention(Q, K, V, d_k) that returns the output matrix. Predict what weights.sum(-1) should be.
Hint: scores = Q @ K.T / d_k**0.5; weights = F.softmax(scores, dim=-1); output = weights @ V. Assert row sums = 1.
import torch, torch.nn.functional as F
import numpy as np
torch.manual_seed(42)
d_k = 4
Q = torch.randn(3, d_k)
K = torch.randn(3, d_k)
V = torch.tensor([[ 0.1,-0.5, 0.8, 0.3],
[ 0.6, 0.2,-0.4, 0.9],
[-0.3, 0.7, 0.1,-0.2]])
def attention(Q, K, V, d_k):
scores = Q @ K.T / d_k**0.5
weights = F.softmax(scores, dim=-1)
return weights @ V, weights
out, W = attention(Q, K, V, d_k)
print('weight row sums:', np.round(W.sum(-1).numpy(), 6))
print('output:', np.round(out.numpy(), 4))| token | W row sums | output[0] |
|---|---|---|
| 0 | 1.000000 | [0.069, 0.259, 0.086, 0.434] |
| 1 | 1.000000 | [0.227, 0.254, -0.009, 0.746] |
| 2 | 1.000000 | [0.053, 0.226, 0.183, 0.290] |
Worked example
Your turn: for each d_k in [4, 16, 64], sample 10,000 (q, k) pairs from N(0,1) and compute (q·k).var(). Predict what value you'll see.
Hint: q=torch.randn(10000, d); k=torch.randn(10000, d); dots=(q*k).sum(1). Print dots.var().item() — it should equal d_k.
import torch
for d in [4, 16, 64]:
q = torch.randn(10000, d)
k = torch.randn(10000, d)
dots = (q * k).sum(1)
print(f'd_k={d:2d}: var(q^T k)={dots.var().item():.2f} '
f'(expected {d}), std={dots.std().item():.3f}')| d_k | Var(q·k) empirical | expected | std |
|---|---|---|---|
| 4 | 3.93 | 4 | 1.983 |
| 16 | 16.10 | 16 | 4.013 |
| 64 | 64.67 | 64 | 8.042 |
Trade off
Comparison matrix
From Milestone 2 — prove Var(qᵀk) = d_k: every row here is a choice with a cost. Fill the Var(q·k) empirical column, then say which row you would actually pick and what you give up for it.
| d_k | Var(q·k) empirical | expected | std |
|---|---|---|---|
| 4 | 3.93 | 4 | 1.983 |
| 16 | 16.10 | 16 | 4.013 |
| 64 | 64.67 | 64 | 8.042 |
Worked example
Your turn: combine the attention function and the saturation demo. Print the entropy for d_k ∈ {4,16,64,256} both with and without scaling.
Hint: entropy = -(p * torch.log(p + 1e-9)).sum(). At d_k=256, unscaled entropy should drop to near 0.
import torch, torch.nn.functional as F
def entropy(logits):
p = F.softmax(logits, dim=-1)
return -(p * torch.log(p + 1e-9)).sum().item()
for d in [4, 16, 64, 256]:
torch.manual_seed(0)
q = torch.randn(1, d); k = torch.randn(5, d)
raw = (q @ k.T)[0]
print(f'd_k={d:3d}: H(unscaled)={entropy(raw):.3f} '
f'H(scaled)={entropy(raw / d**0.5):.3f}')
print('\nAttention formula verified above.')| d_k | H(unscaled) | H(scaled) | verdict |
|---|---|---|---|
| 4 | 1.240 | 1.491 | both spread |
| 16 | 0.626 | 1.439 | unscaled begins to sharpen |
| 64 | 0.404 | 1.367 | unscaled nearly peaked |
| 256 | 0.000 | 1.065 | unscaled saturated; scaled healthy |
Comparison
Comparison matrix
From Milestone 3 — full program with saturation demo: refill the H(unscaled) column from what you know. The rest of the table is as it appeared.
| d_k | H(unscaled) | H(scaled) | verdict |
|---|---|---|---|
| 4 | 1.240 | 1.491 | both spread |
| 16 | 0.626 | 1.439 | unscaled begins to sharpen |
| 64 | 0.404 | 1.367 | unscaled nearly peaked |
| 256 | 0.000 | 1.065 | unscaled saturated; scaled healthy |
Concept
Out loud, slides closed: (1) derive Attention(Q,K,V) from Bahdanau alignment to the matrix formula; (2) prove Var(qᵀk)=d_k and state why that forces the √d_k divisor; (3) explain why attention is a soft lookup table and how it relates to Nadaraya-Watson regression.
Stretch (homework): implement multi-head attention by splitting d_model into h heads of size d_k=d_model/h each, running attention in parallel, and concatenating; show that the output dimension is unchanged. Next: Lesson 80 — multi-head attention, positional encodings, and the full transformer block.
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — The RNN bottleneck & Bahdanau attention · Queries, keys, values, and self-attention · The √d_k scaling: variance proof · Your turn: implement attention. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
| concept | the one thing to remember |
|---|---|
| RNN bottleneck | whole sequence compressed into one h_T; Bahdanau attends to all encoder states instead |
| Q / K / V | Q=what I seek, K=what I advertise, V=what I contribute |
| √d_k scaling | Var(q·k)=d_k; dividing restores unit variance and prevents softmax collapse |
| self-attention | Q=K=V=XW; output i = ∑_j softmax(sim)[i,j] · V[j] |
| soft lookup | attention = differentiable Nadaraya-Watson regression; hard lookup is the 0-temp limit |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.