USAAIO Lesson 80, from Phase 3, in which you implement scaled_dot_product_attention(Q, K, V, mask) from scratch in PyTorch. It analyzes the shapes of Q, K, and V, gives the rationale for the 1/sqrt(d_k) scaling, and covers causal lower-triangular masking and padding masks. It then covers numerical stability, including FP16 overflow and the max-subtraction fix, the O(T^2 * d_k) complexity, and verification against F.scaled_dot_product_attention. All the trace-table values were verified with torch 2.7.1+cpu in June 2026. The lesson runs to 32 slides.
Subject: Machine Learning · 64 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 80 · Phase 3 (Transformers & NLP begins)
Build scaled_dot_product_attention(Q, K, V, mask) from the ground up: shape analysis, the sqrt(d_k) scaling, causal and padding masks, numerical stability in FP16, and O(T²) complexity — the primitive every transformer reuses.
Objectives
(batch, seq, d_k/d_v) and trace the output shape1/sqrt(d_k) is the correct scalemasked_fill before softmaxF.scaled_dot_product_attentionWarm-up
Discussion prompt
Before we open Lesson 80: Scaled Dot-Product Attention from Scratch: without looking back, what was the main idea of Scaled Dot-Product Attention, 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:
RNN sequence bottleneck, Bahdanau alignment scores, self-attention (Q=K=V from the same sequence), derivation of Attention(Q,K,V) = softmax(QK^T/sqrt(d_k))V, variance proof that dot products grow as d_k, and why scaling prevents softmax saturation.
Section
Part 1 of 4
Concept
Attention answers: for each query token, how much should I attend to each key token? The answer is a weighted sum over the values.
\[ \text{Attention}(Q, K, V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V \]
This is the core primitive in every transformer. Multi-head attention (Lesson 81) runs this function h times in parallel — master this one first.
Counterexample
Discussion prompt
Attention answers: for each query token, how much should I attend to each key token? The answer is a weighted sum over the values.
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.
Concept
All three tensors carry a batch dimension. d_k is the query/key dimension (must match for the dot product); d_v is the value dimension (sets the output width).
| tensor | shape | role |
|---|---|---|
| Q (query) | (batch, seq, d_k) | what we are looking for |
| K (key) | (batch, seq, d_k) | what each position offers |
| V (value) | (batch, seq, d_v) | what each position contributes |
| output | (batch, seq, d_v) | weighted sum of values |
The intermediate attention-weight matrix is (batch, seq, seq) — quadratic in sequence length. This is both the expressiveness and the cost.
Discrimination
Sort into buckets
Sort these by shape, from memory, without looking back at QKV shape contract. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Concept
For q, k ~ N(0,1) (unit normal, independent), the dot product q·k has variance d_k. Without scaling, large d_k pushes scores into the softmax saturation zone.
\[ \mathrm{Var}(q \cdot k) = d_k \quad\Rightarrow\quad \mathrm{Var}\!\left(\frac{q\cdot k}{\sqrt{d_k}}\right) = 1 \]
Verified: with d_k=8, B=2, T=5, seed=0, raw scores have variance 6.60; after dividing by sqrt(8)=2.828 the variance drops to 0.82 — near 1. Saturated softmax → vanishing gradients (Lesson 9 callback).
Analogy
Discussion prompt
Explain Why divide by sqrt(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:
For q, k ~ N(0,1) (unit normal, independent), the dot product q·k has variance d_k. Without scaling, large d_k pushes scores into the softmax saturation zone.
Estimation
Predict first
Build scaled_dot_product_attention step by step. Start: Q @ K^T then scale.
Commit before you compute: what does Implement attention: shapes and scores come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: raw variance = 6.5991 → scaled variance = 0.8249 (near 1)
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. Dividing by sqrt(8)=2.828 normalises the dot-product distribution.
Worked example
Build scaled_dot_product_attention step by step. Start: Q @ K^T then scale.
import torch, torch.nn.functional as F
torch.manual_seed(0)
B, T, d_k, d_v = 2, 5, 8, 8
Q = torch.randn(B, T, d_k)
K = torch.randn(B, T, d_k)
V = torch.randn(B, T, d_v)
# Step 1: raw scores (B, T, T)
scores = Q @ K.transpose(-2, -1) # matmul over last two dims
print('raw var:', round(scores.var().item(), 4))
# Step 2: scale
scores = scores / (d_k ** 0.5) # divide every element
print('scaled var:', round(scores.var().item(), 4))raw variance = 6.5991 → scaled variance = 0.8249 (near 1)
Why: Dividing by sqrt(8)=2.828 normalises the dot-product distribution. Softmax on near-unit-variance scores produces well-spread attention weights rather than near-one-hot spikes.
| step | shape | variance |
|---|---|---|
| Q @ K^T (raw scores) | (2, 5, 5) | 6.5991 |
| / sqrt(d_k=8) | (2, 5, 5) | 0.8249 |
| softmax(scores, dim=-1) | (2, 5, 5) | — |
| weights @ V | (2, 5, 8) | — |
Discrimination
Sort into buckets
Sort these by shape, from memory, without looking back at Implement attention: shapes and scores. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Missing information
Discussion prompt
Add softmax over the key dimension (dim=-1) and multiply by V. Verify the output shape.
What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.
Hint: Anything you would have to invent to get started is a thing the problem must supply.
Answer:
softmax(dim=-1) makes each query's distribution over keys sum to 1. The context vector is then a convex combination of the 5 value rows, weighted by relevance.
Worked example
Add softmax over the key dimension (dim=-1) and multiply by V. Verify the output shape.
# Continuing from previous slide
weights = F.softmax(scores, dim=-1) # (B, T, T); rows sum to 1
context = weights @ V # (B, T, d_v)
print('weights[0,0,:]:', [round(x,4) for x in weights[0,0,:].tolist()])
print('weights row sum:', round(weights[0,0,:].sum().item(), 4))
print('context shape: ', context.shape)weights[0,0,:] = [0.0907, 0.0235, 0.0274, 0.1826, 0.6757]; sum = 1.0
Why: softmax(dim=-1) makes each query's distribution over keys sum to 1. The context vector is then a convex combination of the 5 value rows, weighted by relevance.
| key position | attention weight (query 0) |
|---|---|
| 0 | 0.0907 |
| 1 | 0.0235 |
| 2 | 0.0274 |
| 3 | 0.1826 |
| 4 | 0.6757 |
Comparison
Comparison matrix
From Softmax → context output: refill the attention weight (query 0) column from what you know. The rest of the table is as it appeared.
| key position | attention weight (query 0) |
|---|---|
| 0 | 0.0907 |
| 1 | 0.0235 |
| 2 | 0.0274 |
| 3 | 0.1826 |
| 4 | 0.6757 |
Section
Part 2 of 4
Concept
Unmasked attention lets every position see every other position. Two contexts require restricting this:
| mask type | what it blocks | where used |
|---|---|---|
| causal (autoregressive) | future positions (upper triangle) | decoder self-attention (GPT, L82) |
| padding | pad tokens (variable-length batches) | encoder + decoder everywhere |
The mechanism is identical for both: set the blocked score to -inf before softmax. After exp(-inf) = 0, those positions contribute nothing to the weighted sum.
Trade off
Comparison matrix
From Why masking is necessary: every row here is a choice with a cost. Fill the what it blocks column, then say which row you would actually pick and what you give up for it.
| mask type | what it blocks | where used |
|---|---|---|
| causal (autoregressive) | future positions (upper triangle) | decoder self-attention (GPT, L82) |
| padding | pad tokens (variable-length batches) | encoder + decoder everywhere |
Concept
For a sequence of length T, the causal mask is a (T, T) lower-triangular boolean matrix. True at (i, j) means query i is allowed to attend to key j (i.e., j ≤ i).
import torch
T = 5
# True = allowed; False = blocked (will become -inf)
causal_mask = torch.tril(torch.ones(T, T, dtype=torch.bool))
print(causal_mask.int().numpy())| row (query) | allowed key positions | blocked key positions |
|---|---|---|
| 0 | 0 only | 1, 2, 3, 4 |
| 1 | 0, 1 | 2, 3, 4 |
| 2 | 0, 1, 2 | 3, 4 |
| 3 | 0, 1, 2, 3 | 4 |
| 4 | 0, 1, 2, 3, 4 | none |
Pattern
Step through it
Step through Causal mask: lower-triangular boolean one row at a time. What is driving the change, and what would the row after the last one be?
Pattern
Predict first
The table runs: 0 (causal) | 1.0000 | 0.0000 | 0.0000 | 0.0000 | 0.0000 · 2 (causal) | 0.2639 | 0.1116 | 0.6245 | 0.0000 | 0.0000
In Apply causal mask with masked_fill, given the rows so far: what is the next one — the row where query pos is 2 (unmasked)?
Correct: 2 (unmasked) | 0.1005 | 0.0425 | 0.2378 | 0.0937 | 0.5256
| query pos | w[k=0] | w[k=1] | w[k=2] | w[k=3] | w[k=4] |
|---|---|---|---|---|---|
| 0 (causal) | 1.0000 | 0.0000 | 0.0000 | 0.0000 | 0.0000 |
| 2 (causal) | 0.2639 | 0.1116 | 0.6245 | 0.0000 | 0.0000 |
| 2 (unmasked) | 0.1005 | 0.0425 | 0.2378 | 0.0937 | 0.5256 |
Why: The relationship between the columns, not the individual numbers, is what generates the next row. After -inf fill, exp(-inf)=0 exactly.
Worked example
Apply the causal mask to the scaled scores: where ~mask (blocked), replace with -inf. Then softmax collapses those positions to zero weight.
# scores shape: (B, T, T); causal_mask shape: (T, T)
masked_scores = scores.masked_fill(~causal_mask, float('-inf'))
attn_masked = F.softmax(masked_scores, dim=-1)
# Position 0: can only see itself
print('pos 0 weights:', [round(x,4) for x in attn_masked[0,0,:].tolist()])
# Position 2: sees positions 0,1,2
print('pos 2 weights:', [round(x,4) for x in attn_masked[0,2,:].tolist()])pos 0 weights = [1.0, 0.0, 0.0, 0.0, 0.0] — attends only to itself
Why: After -inf fill, exp(-inf)=0 exactly. softmax normalises the remaining positions. Position 0 has no past, so weight 1.0 on itself is correct.
| query pos | w[k=0] | w[k=1] | w[k=2] | w[k=3] | w[k=4] |
|---|---|---|---|---|---|
| 0 (causal) | 1.0000 | 0.0000 | 0.0000 | 0.0000 | 0.0000 |
| 2 (causal) | 0.2639 | 0.1116 | 0.6245 | 0.0000 | 0.0000 |
| 2 (unmasked) | 0.1005 | 0.0425 | 0.2378 | 0.0937 | 0.5256 |
Pattern
Step through it
Step through Apply causal mask with masked_fill one row at a time. What is driving the change, and what would the row after the last one be?
Anomaly
Predict first
A student writes this, and it looks reasonable:
Compute weights = F.softmax(scores, dim=-1), then zero out the forbidden positions: weights[~mask] = 0.0.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Zeroing AFTER softmax leaves weights [0.091, 0, 0, 0, 0] — they no longer sum to 1.
Set blocked scores to -inf BEFORE softmax: scores.masked_fill(~mask, float('-inf')).
Why: Zeroing AFTER softmax leaves weights [0.091, 0, 0, 0, 0] — they no longer sum to 1. The context vector is scaled by 0.091, not 1. Downstream layers receive wrong magnitudes, gradients are corrupted.
Trap
Compute weights = F.softmax(scores, dim=-1), then zero out the forbidden positions: weights[~mask] = 0.0.
weights[0,0,:] = softmax([−0.3979, −1.7482, −1.5943, 0.3019, 1.6101]) = [0.091, 0.024, 0.027, 0.183, 0.676], then zero out positions 1-4
Why: Zeroing AFTER softmax leaves weights [0.091, 0, 0, 0, 0] — they no longer sum to 1. The context vector is scaled by 0.091, not 1. Downstream layers receive wrong magnitudes, gradients are corrupted.
Set blocked scores to -inf BEFORE softmax: scores.masked_fill(~mask, float('-inf')).
After -inf fill, scores[0,0,:] = [−0.3979, -inf, -inf, -inf, -inf]; softmax → [1.0, 0, 0, 0, 0]
Why: exp(-inf)=0 so the forbidden positions vanish before the softmax normalisation. The remaining weights sum exactly to 1 — the context is a proper convex combination of the allowed values.
Break the constraint
Discussion prompt
The rule this trap just fixed:
Set blocked scores to -inf BEFORE softmax: scores.masked_fill(~mask, float('-inf')).
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:
Zeroing AFTER softmax leaves weights [0.091, 0, 0, 0, 0] — they no longer sum to 1. The context vector is scaled by 0.091, not 1. Downstream layers receive wrong magnitudes, gradients are corrupted.
Section
Part 3 of 4
Concept
FP16 (half-precision) has maximum representable value ≈ 65504. For attention scores with typical values of 10-50 (long sequences, large d_model), exp(score) overflows to inf before normalisation.
Verified: torch.exp(torch.tensor(20.0).half()) = inf. With big_scores = [10, 20, 30, 40, 50]: FP16 exp values = [22032, inf, inf, inf, inf]. Softmax of inf/inf = nan.
| precision | exp(20) | exp(30) | exp(50) | softmax safe? |
|---|---|---|---|---|
| FP32 | 4.85e+08 | 1.07e+13 | 5.18e+21 | yes (large range) |
| FP16 | inf | inf | inf | nan / corrupted |
Explain it
Discussion prompt
Explain FP16 overflow in naive softmax to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.
Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.
Answer:
FP16 (half-precision) has maximum representable value ≈ 65504. For attention scores with typical values of 10-50 (long sequences, large d_model), exp(score) overflows to inf before normalisation.
Concept
Subtract the row maximum before exp. This is mathematically equivalent (the max cancels in numerator and denominator) but keeps all exponent arguments ≤ 0 — so exp never overflows.
\[ \text{softmax}(x_i) = \frac{e^{x_i - m}}{\sum_j e^{x_j - m}} \quad\text{where } m = \max_j x_j \]
PyTorch's F.softmax already does this internally. In a hand-rolled implementation you must do it explicitly. F.scaled_dot_product_attention (PyTorch 2.0+) does it for you — verified match: max abs diff = 2.38e-07 vs our scratch version.
Analogy
Discussion prompt
Explain The max-subtraction fix 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:
Subtract the row maximum before exp. This is mathematically equivalent (the max cancels in numerator and denominator) but keeps all exponent arguments ≤ 0 — so exp never overflows.
Estimation
Predict first
Assemble the complete function with optional mask and the safe softmax path.
Commit before you compute: what does Full scratch implementation + stability come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: match = True, max abs diff = 2.38e-07
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. Our implementation is numerically identical to PyTorch's fused kernel (within float32 rounding).
Worked example
Assemble the complete function with optional mask and the safe softmax path.
import torch, torch.nn.functional as F
def scaled_dot_product_attention(Q, K, V, mask=None):
d_k = Q.size(-1)
scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(~mask, float('-inf'))
weights = F.softmax(scores, dim=-1) # stable internally
return weights @ V, weights
torch.manual_seed(0)
B, T, d_k = 2, 5, 8
Q = torch.randn(B, T, d_k)
K = torch.randn(B, T, d_k)
V = torch.randn(B, T, d_k)
out, w = scaled_dot_product_attention(Q, K, V)
torch_ref = F.scaled_dot_product_attention(Q, K, V)
print('match:', torch.allclose(out, torch_ref, atol=1e-5))
print('max diff:', f'{(out-torch_ref).abs().max().item():.2e}')match = True, max abs diff = 2.38e-07
Why: Our implementation is numerically identical to PyTorch's fused kernel (within float32 rounding). This validates both the formula and the stable-softmax path.
| implementation | output shape | matches torch? | max diff |
|---|---|---|---|
| scratch (no mask) | (2,5,8) | True | 2.38e-07 |
| scratch (causal mask) | (2,5,8) | True | 1.19e-07 |
| F.scaled_dot_product_attention | (2,5,8) | reference | — |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
match = True, max abs diff = 2.38e-07
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:
Assemble the complete function with optional mask and the safe softmax path.
Section
Part 4 of 4
Concept
The bottleneck is Q @ K^T: a (T, d_k) @ (d_k, T) matmul that costs O(T² · d_k) FLOPs and produces an O(T²) attention matrix. Memory is dominated by that matrix.
| seq len T | d_k=64 | QK^T FLOPs | attn matrix entries |
|---|---|---|---|
| 128 | 64 | 1,048,576 | 16,384 |
| 512 | 64 | 16,777,216 | 262,144 |
| 2048 | 64 | 268,435,456 | 4,194,304 |
This is why naive attention is quadratic in context length — and why FlashAttention (Lesson 90+) rewrites the tiling to avoid materialising the full (T, T) matrix in HBM.
Comparison
Comparison matrix
From Computational cost: O(T² · d_k): refill the attn matrix entries column from what you know. The rest of the table is as it appeared.
| seq len T | d_k=64 | QK^T FLOPs | attn matrix entries |
|---|---|---|---|
| 128 | 64 | 1,048,576 | 16,384 |
| 512 | 64 | 16,777,216 | 262,144 |
| 2048 | 64 | 268,435,456 | 4,194,304 |
Missing information
Discussion prompt
Walk every number by hand with T=3, d_k=4, d_v=2. Fixed Q3, K3, V3 (no randomness — exact reproducible values).
What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.
Hint: Anything you would have to invent to get started is a thing the problem must supply.
Answer:
With d_k=4, sqrt(d_k)=2. Q3[0] @ K3^T = [0.75, -1.0, -0.75]; divide by 2 gives [0.375, -0.5, -0.375]; softmax gives the weight row above.
Worked example
Walk every number by hand with T=3, d_k=4, d_v=2. Fixed Q3, K3, V3 (no randomness — exact reproducible values).
import torch, torch.nn.functional as F
Q3 = torch.tensor([[ 1.0, 0.0,-1.0, 0.5],
[ 0.5, 1.0, 0.0,-0.5],
[-0.5, 0.5, 1.0, 0.0]])
K3 = torch.tensor([[ 1.0, 0.5, 0.0,-0.5],
[-0.5, 1.0, 0.5, 0.0],
[ 0.0,-0.5, 1.0, 0.5]])
V3 = torch.tensor([[1.0, 0.0],[0.0, 1.0],[0.5, 0.5]])
scale = 4 ** 0.5 # sqrt(d_k=4) = 2.0
scores = Q3 @ K3.t() / scale
weights = F.softmax(scores, dim=-1)
out = weights @ V3
print('scores:'); print(scores.numpy().round(4))
print('weights:'); print(weights.numpy().round(4))
print('output:'); print(out.numpy().round(4))scale = sqrt(4) = 2.0; scores[0,:] = [0.375, -0.5, -0.375]; weights[0,:] = [0.5293, 0.2207, 0.2500]
Why: With d_k=4, sqrt(d_k)=2. Q3[0] @ K3^T = [0.75, -1.0, -0.75]; divide by 2 gives [0.375, -0.5, -0.375]; softmax gives the weight row above.
| query | scores (/ sqrt(4)=2) | weights (softmax) | output |
|---|---|---|---|
| q0 | [0.375, -0.500, -0.375] | [0.5293, 0.2207, 0.2500] | [0.6543, 0.3457] |
| q1 | [0.625, 0.375, -0.375] | [0.4658, 0.3628, 0.1714] | [0.5515, 0.4485] |
| q2 | [-0.125, 0.625, 0.375] | [0.2098, 0.4442, 0.3460] | [0.3828, 0.6172] |
Estimation
Predict first
Apply the causal mask to the same Q3/K3/V3 and trace how query 1's weights redistribute.
Commit before you compute: what does Causal mask on the 3-token example come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: q0 attends only to k0 (weight=1.0); q1 attends to k0 (0.5622) and k1 (0.4378); q2 unaffected (same as unmasked)
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. q2 is the last token — it can see all 3 keys, so masking changes nothing for it.
Worked example
Apply the causal mask to the same Q3/K3/V3 and trace how query 1's weights redistribute.
# Continuing with Q3, K3, V3, scores from previous slide
cm3 = torch.tril(torch.ones(3, 3, dtype=torch.bool))
masked_scores = scores.masked_fill(~cm3, float('-inf'))
w_causal = F.softmax(masked_scores, dim=-1)
out_causal = w_causal @ V3
print('causal weights:'); print(w_causal.numpy().round(4))
print('causal output:'); print(out_causal.numpy().round(4))q0 attends only to k0 (weight=1.0); q1 attends to k0 (0.5622) and k1 (0.4378); q2 unaffected (same as unmasked)
Why: q2 is the last token — it can see all 3 keys, so masking changes nothing for it. q1 loses k2 (-inf), so its remaining weight normalises between k0 and k1 only.
| query | w[k=0] | w[k=1] | w[k=2] | output |
|---|---|---|---|---|
| q0 (causal) | 1.0000 | 0.0000 | 0.0000 | [1.0000, 0.0000] |
| q1 (causal) | 0.5622 | 0.4378 | 0.0000 | [0.5622, 0.4378] |
| q2 (causal = unmasked) | 0.2098 | 0.4442 | 0.3460 | [0.3828, 0.6172] |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
q0 attends only to k0 (weight=1.0); q1 attends to k0 (0.5622) and k1 (0.4378); q2 unaffected (same as unmasked)
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:
Apply the causal mask to the same Q3/K3/V3 and trace how query 1's weights redistribute.
Ranking
Put in order
These are the steps of The scaled dot-product attention recipe, scrambled. Put them back in order before the next slide shows you.
(batch, seq, d_k), V (batch, seq, d_v) → output (batch, seq, d_v)Q @ K.transpose(-2,-1) → (batch, seq, seq); divide by sqrt(d_k)scores.masked_fill(~mask, float('-inf')) — always BEFORE softmaxF.softmax(scores, dim=-1) — safe against overflow, rows sum to 1weights @ V — weighted sum of values; verify against F.scaled_dot_product_attentionWhy: 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
(batch, seq, d_k), V (batch, seq, d_v) → output (batch, seq, d_v)Q @ K.transpose(-2,-1) → (batch, seq, seq); divide by sqrt(d_k)scores.masked_fill(~mask, float('-inf')) — always BEFORE softmaxF.softmax(scores, dim=-1) — safe against overflow, rows sum to 1weights @ V — weighted sum of values; verify against F.scaled_dot_product_attentionEdge cases
Discussion prompt
The scaled dot-product attention 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:
(batch, seq, d_k), V (batch, seq, d_v) → output (batch, seq, d_v)Q @ K.transpose(-2,-1) → (batch, seq, seq); divide by sqrt(d_k)scores.masked_fill(~mask, float('-inf')) — always BEFORE softmaxF.softmax(scores, dim=-1) — safe against overflow, rows sum to 1weights @ V — weighted sum of values; verify against F.scaled_dot_product_attentionElimination
Eliminate the wrong options
Q has shape (4, 10, 64), K has shape (4, 10, 64), V has shape (4, 10, 128). What is the shape of the attention output?
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: Output shape = (batch, seq, d_v). The attention-weight matrix is (4,10,10); multiplying by V of shape (4,10,128) gives (4,10,128). d_v=128 sets the output width, not d_k=64.
Check
Trace the shapes before clicking.
Check your understanding
Q has shape (4, 10, 64), K has shape (4, 10, 64), V has shape (4, 10, 128). What is the shape of the attention output?
Answer: A
Why: Output shape = (batch, seq, d_v). The attention-weight matrix is (4,10,10); multiplying by V of shape (4,10,128) gives (4,10,128). d_v=128 sets the output width, not d_k=64.
Prediction
Predict first
For d_k=64 (a typical transformer head size), what is the correct scale divisor, and why 1/sqrt(d_k) specifically?
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: 8.0 — normalises dot-product variance to ~1 for unit-normal Q, K
Why: For q, k ~ N(0,1), Var(q·k) = d_k, so the standard deviation is sqrt(d_k) = sqrt(64) = 8. Dividing by sqrt(d_k) gives Var = 1. Verified: raw variance 6.60 → scaled variance 0.82 at d_k=8.
Check
Reason from first principles.
Check your understanding
For d_k=64 (a typical transformer head size), what is the correct scale divisor, and why 1/sqrt(d_k) specifically?
Answer: A
Why: For q, k ~ N(0,1), Var(q·k) = d_k, so the standard deviation is sqrt(d_k) = sqrt(64) = 8. Dividing by sqrt(d_k) gives Var = 1. Verified: raw variance 6.60 → scaled variance 0.82 at d_k=8.
Elimination
Eliminate the wrong options
A student writes: weights = F.softmax(scores, dim=-1); weights[~causal_mask] = 0.0. What is wrong?
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: After softmax, a row like [0.09, 0.02, 0.03, 0.18, 0.68] sums to 1.0. Zeroing the last 4 leaves [0.09, 0, 0, 0, 0] — sum = 0.09, not 1. The context vector is scaled by 0.09 instead of 1, corrupting downstream projections and gradients. The -inf-before-softmax trick avoids this because exp(-inf)=0 before normalisation.
Check
Identify the bug.
Check your understanding
A student writes: weights = F.softmax(scores, dim=-1); weights[~causal_mask] = 0.0. What is wrong?
Answer: A
Why: After softmax, a row like [0.09, 0.02, 0.03, 0.18, 0.68] sums to 1.0. Zeroing the last 4 leaves [0.09, 0, 0, 0, 0] — sum = 0.09, not 1. The context vector is scaled by 0.09 instead of 1, corrupting downstream projections and gradients. The -inf-before-softmax trick avoids this because exp(-inf)=0 before normalisation.
Section
Project
Concept
Implement scaled_dot_product_attention(Q, K, V, mask=None), test on the 3-token example, add a causal mask, and verify against F.scaled_dot_product_attention.
| # | milestone | key tool |
|---|---|---|
| 1 | Compute scaled scores; match variance table | @ operator, .transpose(-2,-1) |
| 2 | Add softmax → output; verify 3-token trace numbers | F.softmax(dim=-1) |
| 3 | Causal mask: verify q0 gets weight 1.0 on itself | torch.tril, masked_fill |
| 4 | Verify vs F.scaled_dot_product_attention (match=True) | torch.allclose, atol=1e-5 |
Build rules: type every line; check shapes after each operation with print(x.shape); and do NOT apply the mask after softmax — always before.
Counterexample
Discussion prompt
Implement scaled_dot_product_attention(Q, K, V, mask=None), test on the 3-token example, add a causal mask, and verify against F.scaled_dot_product_attention.
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; check shapes after each operation with print(x.shape); and do NOT apply the mask after softmax — always before.
Worked example
Your turn: compute scaled scores and softmax weights for the 3-token example. Predict the weight for query 0 attending to key 0.
Hint: scores = Q3 @ K3.t() / (d_k ** 0.5) then F.softmax(scores, dim=-1). With d_k=4, scale=2.0. Expected q0 weights: [0.5293, 0.2207, 0.2500].
import torch, torch.nn.functional as F
Q3 = torch.tensor([[ 1.0, 0.0,-1.0, 0.5],
[ 0.5, 1.0, 0.0,-0.5],
[-0.5, 0.5, 1.0, 0.0]])
K3 = torch.tensor([[ 1.0, 0.5, 0.0,-0.5],
[-0.5, 1.0, 0.5, 0.0],
[ 0.0,-0.5, 1.0, 0.5]])
V3 = torch.tensor([[1.0,0.0],[0.0,1.0],[0.5,0.5]])
d_k = Q3.size(-1) # 4
scores = Q3 @ K3.t() / (d_k**0.5) # (3,3)
weights = F.softmax(scores, dim=-1) # rows sum to 1
print('weights:', weights.numpy().round(4))
print('output: ', (weights @ V3).numpy().round(4))| query | w[k=0] | w[k=1] | w[k=2] | output[0] | output[1] |
|---|---|---|---|---|---|
| q0 | 0.5293 | 0.2207 | 0.2500 | 0.6543 | 0.3457 |
| q1 | 0.4658 | 0.3628 | 0.1714 | 0.5515 | 0.4485 |
| q2 | 0.2098 | 0.4442 | 0.3460 | 0.3828 | 0.6172 |
Pattern
Step through it
Step through Milestone 1 & 2 — scores and weights one row at a time. What is driving the change, and what would the row after the last one be?
Worked example
Your turn: add a causal mask and verify the full function against F.scaled_dot_product_attention.
Hint: torch.tril(torch.ones(T, T, dtype=torch.bool)); apply with masked_fill(~mask, float('-inf')) before softmax. Then test torch.allclose(out, F.scaled_dot_product_attention(Q, K, V), atol=1e-5).
def scaled_dot_product_attention(Q, K, V, mask=None):
d_k = Q.size(-1)
scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(~mask, float('-inf'))
return F.softmax(scores, dim=-1) @ V
# Test causal mask
cm = torch.tril(torch.ones(3, 3, dtype=torch.bool))
out_c = scaled_dot_product_attention(Q3, K3, V3, mask=cm)
print('q0 output (causal):', out_c[0].numpy().round(4))
# Verify no-mask path
torch.manual_seed(0)
Q2=torch.randn(2,5,8); K2=torch.randn(2,5,8); V2=torch.randn(2,5,8)
ref = F.scaled_dot_product_attention(Q2, K2, V2)
print('match:', torch.allclose(scaled_dot_product_attention(Q2,K2,V2), ref, atol=1e-5))| test | result |
|---|---|
| q0 output causal | [1.0000, 0.0000] (attends only to v0=[1,0]) |
| q1 output causal | [0.5622, 0.4378] |
| match vs F.scaled_dot_product_attention | True |
Trade off
Comparison matrix
From Milestone 3 & 4 — causal mask + verification: every row here is a choice with a cost. Fill the result column, then say which row you would actually pick and what you give up for it.
| test | result |
|---|---|
| q0 output causal | [1.0000, 0.0000] (attends only to v0=[1,0]) |
| q1 output causal | [0.5622, 0.4378] |
| match vs F.scaled_dot_product_attention | True |
Concept
import torch, torch.nn.functional as F
def scaled_dot_product_attention(Q, K, V, mask=None):
"""Q,K: (B,T,d_k) V: (B,T,d_v) mask: (T,T) bool (True=allowed)"""
d_k = Q.size(-1)
scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(~mask, float('-inf'))
weights = F.softmax(scores, dim=-1)
return weights @ V, weights
torch.manual_seed(0)
B, T, d_k = 2, 5, 8
Q = torch.randn(B, T, d_k)
K = torch.randn(B, T, d_k)
V = torch.randn(B, T, d_k)
# No mask
out, w = scaled_dot_product_attention(Q, K, V)
ref = F.scaled_dot_product_attention(Q, K, V)
print('no-mask match:', torch.allclose(out, ref, atol=1e-5))
# Causal mask
mask = torch.tril(torch.ones(T, T, dtype=torch.bool))
out_c, w_c = scaled_dot_product_attention(Q, K, V, mask=mask)
ref_c = F.scaled_dot_product_attention(Q, K, V, is_causal=True)
print('causal match: ', torch.allclose(out_c, ref_c, atol=1e-4))| path | match vs F.scaled_dot_product_attention | max abs diff |
|---|---|---|
| no mask | True | 2.38e-07 |
| causal mask | True | 1.19e-07 |
If both lines print True — you have built the primitive every transformer in Phase 3 reuses. Multi-head attention (Lesson 81) is just this function called h times with projected Q, K, V.
Comparison
Comparison matrix
From The full program: refill the match vs F.scaled_dot_product_attention column from what you know. The rest of the table is as it appeared.
| path | match vs F.scaled_dot_product_attention | max abs diff |
|---|---|---|
| no mask | True | 2.38e-07 |
| causal mask | True | 1.19e-07 |
Concept
Out loud, slides closed: (1) state the QKV shape contract and derive the output shape; (2) explain why the scale is 1/sqrt(d_k), not 1/d_k; (3) walk through why masking must happen before softmax, not after.
Stretch (homework per lesson plan): implement a causal_mask as a boolean mask and verify each position only sees past positions. Experiment: print attention scores before vs after masking for a long sequence and confirm they differ significantly. Next up: Lesson 81 — Multi-Head Attention (run this h times in parallel with projected Q, K, V).
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — QKV shapes and the attention formula · Masking: causal and padding · Numerical stability in FP16 · Complexity and the 3-token trace · Your turn: implement it. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
(batch, seq, d_v)scaled_dot_product_attention from scratch (scores → scale → mask → softmax → weighted sum)masked_fill BEFORE softmaxF.scaled_dot_product_attention| concept | the one thing to remember |
|---|---|
| scale factor | 1/sqrt(d_k) — normalises dot-product variance to ~1 |
| mask application | masked_fill(-inf) BEFORE softmax, never zero AFTER |
| causal mask | torch.tril(ones(T,T,bool)) — lower triangle = allowed |
| FP16 stability | subtract row max before exp; F.softmax does this for you |
| complexity | O(T²·d_k) time, O(T²) memory — quadratic in seq length |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.