Lesson 80: Scaled Dot-Product Attention from Scratch

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

What this lesson covers

The lesson, slide by slide

1. Scaled Dot-Product Attention from Scratch

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.

2. By the end of this lesson you can

Objectives

  1. State the QKV shape contract (batch, seq, d_k/d_v) and trace the output shape
  2. Implement scaled dot-product attention from scratch and explain why 1/sqrt(d_k) is the correct scale
  3. Construct a causal (lower-triangular) mask and apply it correctly with masked_fill before softmax
  4. Explain the FP16 overflow risk and fix it with the max-subtraction numerically-stable softmax
  5. State the O(T² · d_k) time and O(T²) memory cost, and verify your implementation against F.scaled_dot_product_attention

3. What survived from Scaled Dot-Product Attention?

Warm-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.

4. QKV shapes and the attention formula

Section

Part 1 of 4

5. What is scaled dot-product attention?

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.

6. Break it if you can: What is scaled dot-product attention?

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.

7. QKV shape contract

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).

tensorshaperole
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.

8. Which is which, by shape

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.

(batch, seq, d_k)
Q (query); K (key)
(batch, seq, d_v)
V (value); output
g1
shape is "(batch, seq, d_k)" for Q (query), K (key) — that is what the table on "QKV shape contract" records, and it is the single property separating this group from the rest.
g2
shape is "(batch, seq, d_v)" for V (value), output — that is what the table on "QKV shape contract" records, and it is the single property separating this group from the rest.

9. Why divide by sqrt(d_k)?

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).

10. By analogy: Why divide by sqrt(d_k)?

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.

11. Guess the shape of the answer: Implement attention: shapes and scores

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.

12. Implement attention: shapes and scores

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.

stepshapevariance
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)—

13. Which is which, by shape

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.

(2, 5, 5)
Q @ K^T (raw scores); / sqrt(d_k=8); softmax(scores, dim=-1)
(2, 5, 8)
weights @ V
g1
shape is "(2, 5, 5)" for Q @ K^T (raw scores), / sqrt(d_k=8), softmax(scores, dim=-1) — that is what the table on "Implement attention: shapes and scores" records, and it is the single property separating this group from the rest.
g2
shape is "(2, 5, 8)" for weights @ V — that is what the table on "Implement attention: shapes and scores" records, and it is the single property separating this group from the rest.

14. What has to be given first: Softmax → context output

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.

15. Softmax → context output

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 positionattention weight (query 0)
00.0907
10.0235
20.0274
30.1826
40.6757

16. Fill in: attention weight (query 0) for Softmax → context output

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 positionattention weight (query 0)
00.0907
10.0235
20.0274
30.1826
40.6757

17. Masking: causal and padding

Section

Part 2 of 4

18. Why masking is necessary

Concept

Unmasked attention lets every position see every other position. Two contexts require restricting this:

mask typewhat it blockswhere used
causal (autoregressive)future positions (upper triangle)decoder self-attention (GPT, L82)
paddingpad 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.

19. What each one costs: Why masking is necessary

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 typewhat it blockswhere used
causal (autoregressive)future positions (upper triangle)decoder self-attention (GPT, L82)
paddingpad tokens (variable-length batches)encoder + decoder everywhere

20. Causal mask: lower-triangular boolean

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 positionsblocked key positions
00 only1, 2, 3, 4
10, 12, 3, 4
20, 1, 23, 4
30, 1, 2, 34
40, 1, 2, 3, 4none

21. Watch it run: Causal mask: lower-triangular boolean

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?

  1. Step 1: row (query) is 0
  2. Step 2: row (query) is 1
  3. Step 3: row (query) is 2
  4. Step 4: row (query) is 3
  5. Step 5: row (query) is 4

22. Predict the next row: Apply causal mask with masked_fill

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 posw[k=0]w[k=1]w[k=2]w[k=3]w[k=4]
0 (causal)1.00000.00000.00000.00000.0000
2 (causal)0.26390.11160.62450.00000.0000
2 (unmasked)0.10050.04250.23780.09370.5256

Why: The relationship between the columns, not the individual numbers, is what generates the next row. After -inf fill, exp(-inf)=0 exactly.

23. Apply causal mask with masked_fill

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 posw[k=0]w[k=1]w[k=2]w[k=3]w[k=4]
0 (causal)1.00000.00000.00000.00000.0000
2 (causal)0.26390.11160.62450.00000.0000
2 (unmasked)0.10050.04250.23780.09370.5256

24. Watch it run: Apply causal mask with masked_fill

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?

  1. Step 1: query pos is 0 (causal)
  2. Step 2: query pos is 2 (causal)
  3. Step 3: query pos is 2 (unmasked)

25. Something is wrong here: applying the mask AFTER softmax

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.

26. Trap: applying the mask AFTER softmax

Trap

The 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.

The fix

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.

27. Break it on purpose: applying the mask AFTER softmax

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.

28. Numerical stability in FP16

Section

Part 3 of 4

29. FP16 overflow in naive softmax

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.

precisionexp(20)exp(30)exp(50)softmax safe?
FP324.85e+081.07e+135.18e+21yes (large range)
FP16infinfinfnan / corrupted

30. Teach it back: FP16 overflow in naive softmax

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.

31. The max-subtraction fix

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.

32. By analogy: The max-subtraction fix

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.

33. Guess the shape of the answer: Full scratch implementation + stability

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).

34. Full scratch implementation + stability

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.

implementationoutput shapematches torch?max diff
scratch (no mask)(2,5,8)True2.38e-07
scratch (causal mask)(2,5,8)True1.19e-07
F.scaled_dot_product_attention(2,5,8)reference—

35. Work backwards from the answer: Full scratch implementation + stability

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.

36. Complexity and the 3-token trace

Section

Part 4 of 4

37. Computational cost: O(T² · d_k)

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 Td_k=64QK^T FLOPsattn matrix entries
128641,048,57616,384
5126416,777,216262,144
204864268,435,4564,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.

38. Fill in: attn matrix entries for Computational cost: O(T² · d_k)

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 Td_k=64QK^T FLOPsattn matrix entries
128641,048,57616,384
5126416,777,216262,144
204864268,435,4564,194,304

39. What has to be given first: 3-token trace table (d_k=4, d_v=2)

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.

40. 3-token trace table (d_k=4, d_v=2)

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.

queryscores (/ 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]

41. Guess the shape of the answer: Causal mask on the 3-token example

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.

42. Causal mask on the 3-token example

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.

queryw[k=0]w[k=1]w[k=2]output
q0 (causal)1.00000.00000.0000[1.0000, 0.0000]
q1 (causal)0.56220.43780.0000[0.5622, 0.4378]
q2 (causal = unmasked)0.20980.44420.3460[0.3828, 0.6172]

43. Work backwards from the answer: Causal mask on the 3-token example

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.

44. Rebuild the recipe: The scaled dot-product attention recipe

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.

  1. Shapes: Q/K (batch, seq, d_k), V (batch, seq, d_v) → output (batch, seq, d_v)
  2. Scores: Q @ K.transpose(-2,-1) → (batch, seq, seq); divide by sqrt(d_k)
  3. Mask (if any): scores.masked_fill(~mask, float('-inf')) — always BEFORE softmax
  4. Weights: F.softmax(scores, dim=-1) — safe against overflow, rows sum to 1
  5. Output: weights @ V — weighted sum of values; verify against F.scaled_dot_product_attention

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.

45. The scaled dot-product attention recipe

Pattern

  1. Shapes: Q/K (batch, seq, d_k), V (batch, seq, d_v) → output (batch, seq, d_v)
  2. Scores: Q @ K.transpose(-2,-1) → (batch, seq, seq); divide by sqrt(d_k)
  3. Mask (if any): scores.masked_fill(~mask, float('-inf')) — always BEFORE softmax
  4. Weights: F.softmax(scores, dim=-1) — safe against overflow, rows sum to 1
  5. Output: weights @ V — weighted sum of values; verify against F.scaled_dot_product_attention

46. Where does it stop working: The scaled dot-product attention recipe

Edge 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:

  1. Shapes: Q/K (batch, seq, d_k), V (batch, seq, d_v) → output (batch, seq, d_v)
  2. Scores: Q @ K.transpose(-2,-1) → (batch, seq, seq); divide by sqrt(d_k)
  3. Mask (if any): scores.masked_fill(~mask, float('-inf')) — always BEFORE softmax
  4. Weights: F.softmax(scores, dim=-1) — safe against overflow, rows sum to 1
  5. Output: weights @ V — weighted sum of values; verify against F.scaled_dot_product_attention

47. Rule out three: Check yourself — output shape

Elimination

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.

  • A. (4, 10, 128)
  • B. (4, 10, 64)
  • C. (4, 64, 128)
  • D. (4, 10, 10)

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.

48. Check yourself — output shape

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?

  • A. (4, 10, 128) (correct)
  • B. (4, 10, 64)
  • C. (4, 64, 128)
  • D. (4, 10, 10)

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.

Why B tempts people
(4,10,64) confuses d_k (the key/query dimension used for dot products) with d_v (the value dimension that determines output width). They are equal in single-head standard attention but differ in many architectures.
Why C tempts people
(4,64,128) transposes the seq and d_k axes — there is no operation in SDPA that produces a (d_k, d_v) output. The seq dimension is always preserved.
Why D tempts people
(4,10,10) is the shape of the intermediate attention-weight matrix (before multiplying by V) — it is not the final output.

49. Answer it before you see the options: Check yourself — the scaling factor

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.

50. Check yourself — the scaling factor

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?

  • A. 8.0 — normalises dot-product variance to ~1 for unit-normal Q, K (correct)
  • B. 64.0 — divides by d_k so each score is a per-element average
  • C. 0.5 — a fixed constant independent of d_k
  • D. 1.0 — no scaling is needed; softmax is scale-invariant

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.

Why B tempts people
Dividing by d_k (not sqrt(d_k)) over-shrinks the scores: variance becomes 1/d_k instead of 1, pushing softmax toward uniform distribution and reducing expressiveness.
Why C tempts people
A fixed constant like 0.5 does not adapt to d_k. As d_k grows the scores still saturate softmax — the scaling must be 1/sqrt(d_k) to maintain unit variance regardless of head size.
Why D tempts people
Softmax is NOT scale-invariant in practice: large inputs saturate it (gradients near zero). The variance of dot products grows with d_k, so large d_k without scaling collapses attention weights to near-one-hot spikes.

51. Rule out three: Check yourself — mask application

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.

  • A. The weights no longer sum to 1 — the context vector has incorrect magnitude
  • B. Setting to 0.0 is identical to setting the score to -inf before softmax
  • C. causal_mask should be the upper triangle, not the lower triangle
  • D. Nothing — zeroing after softmax and -inf before softmax are mathematically equivalent

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.

52. Check yourself — mask application

Check

Identify the bug.

Check your understanding

A student writes: weights = F.softmax(scores, dim=-1); weights[~causal_mask] = 0.0. What is wrong?

  • A. The weights no longer sum to 1 — the context vector has incorrect magnitude (correct)
  • B. Setting to 0.0 is identical to setting the score to -inf before softmax
  • C. causal_mask should be the upper triangle, not the lower triangle
  • D. Nothing — zeroing after softmax and -inf before softmax are mathematically equivalent

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.

Why B tempts people
They are NOT equivalent: -inf before softmax sets exp to 0 before the denominator is computed, so normalisation corrects for it. Zeroing after softmax skips the renormalisation step entirely.
Why C tempts people
The lower-triangular mask is correct for causal attention — it allows each position to attend to itself and past positions (i.e., j ≤ i). The upper triangle is what gets blocked.
Why D tempts people
They are not equivalent — see answer A. This is one of the most common attention implementation bugs.

53. Your turn: implement it

Section

Project

54. Project: scaled_dot_product_attention from scratch

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.

#milestonekey tool
1Compute scaled scores; match variance table@ operator, .transpose(-2,-1)
2Add softmax → output; verify 3-token trace numbersF.softmax(dim=-1)
3Causal mask: verify q0 gets weight 1.0 on itselftorch.tril, masked_fill
4Verify 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.

55. Break it if you can: Project: scaled_dot_product_attention from scratch

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.

56. Milestone 1 & 2 — scores and weights

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))
queryw[k=0]w[k=1]w[k=2]output[0]output[1]
q00.52930.22070.25000.65430.3457
q10.46580.36280.17140.55150.4485
q20.20980.44420.34600.38280.6172

57. Watch it run: Milestone 1 & 2 — scores and weights

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?

  1. Step 1: query is q0
  2. Step 2: query is q1
  3. Step 3: query is q2

58. Milestone 3 & 4 — causal mask + verification

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))
testresult
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_attentionTrue

59. What each one costs: Milestone 3 & 4 — causal mask + verification

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.

testresult
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_attentionTrue

60. The full program

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))
pathmatch vs F.scaled_dot_product_attentionmax abs diff
no maskTrue2.38e-07
causal maskTrue1.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.

61. Fill in: match vs F.scaled_dot_product_attention for The full program

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.

pathmatch vs F.scaled_dot_product_attentionmax abs diff
no maskTrue2.38e-07
causal maskTrue1.19e-07

62. Show it off

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).

63. Connect it up: Lesson 80: Scaled Dot-Product Attention from Scratch

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.

64. What you can do now

Recap

conceptthe one thing to remember
scale factor1/sqrt(d_k) — normalises dot-product variance to ~1
mask applicationmasked_fill(-inf) BEFORE softmax, never zero AFTER
causal masktorch.tril(ones(T,T,bool)) — lower triangle = allowed
FP16 stabilitysubtract row max before exp; F.softmax does this for you
complexityO(T²·d_k) time, O(T²) memory — quadratic in seq length

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 80 — Scaled Dot-Product Attention — Barron · USAAIO Round 2 Preparation, 2026
  2. Attention Is All You Need (Vaswani et al., 2017) — Section 3.2.1 Scaled Dot-Product Attention — arXiv:1706.03762
  3. All trace-table values, masking outputs, softmax weights, and F.scaled_dot_product_attention match verified with torch 2.7.1+cpu and numpy 2.2.6 — Real execution, June 2026

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108