USAAIO Lesson 66, from Week 23 of Phase 3. It compares the numerical ranges of FP16 and FP32, then covers mixed-precision training with torch.autocast and GradScaler, loss scaling to prevent FP16 underflow, gradient accumulation as a way to simulate a large batch, and the arithmetic of the memory budget. You build an AMP training loop and verify it against full precision. The lesson runs to 29 slides.
Subject: Machine Learning · 58 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 66 · Week 23 (Phase 3)
FP16 arithmetic, GradScaler, gradient accumulation. Train faster and on larger models — without changing a single hyperparameter.
Objectives
autocast → scaler.scale → scaler.step → scaler.updateWarm-up
Discussion prompt
Before we open Lesson 66: Mixed Precision, Loss Scaling & Gradient Accumulation: without looking back, what was the main idea of Naive Bayes Classifiers, 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:
Multinomial, Bernoulli, and Complement Naive Bayes from first principles — Laplace smoothing derivation, log-space arithmetic to prevent underflow, MNB vs BNB on count vs binary features, ComplementNB for imbalanced short texts, and the proof that NB is a linear classifier in log-space. Implement MultinomialNB from scratch and benchmark against logistic regression.
Section
Part 1 of 4
Concept
FP32 and FP16 differ in how many bits they use for the exponent and mantissa — that gap determines both range and precision.
| format | bits | max value | smallest normal | machine eps |
|---|---|---|---|---|
| FP32 | 32 | 3.40e+38 | 1.18e-38 | 1.19e-07 |
| FP16 | 16 | 6.55e+04 | 6.10e-05 | 9.77e-04 |
FP16's max is ~65 000; anything larger becomes inf (overflow). FP16's smallest representable value is ~6e-5; anything smaller becomes 0 (underflow).
Comparison
Comparison matrix
From Two floating-point formats: refill the bits column from what you know. The rest of the table is as it appeared.
| format | bits | max value | smallest normal | machine eps |
|---|---|---|---|---|
| FP32 | 32 | 3.40e+38 | 1.18e-38 | 1.19e-07 |
| FP16 | 16 | 6.55e+04 | 6.10e-05 | 9.77e-04 |
Concept
The idea: run the forward pass and backward pass in FP16 (fast, half the bandwidth), but keep optimizer state and master weights in FP32 (numerically stable).
| tensor | dtype | reason |
|---|---|---|
| activations (fwd) | FP16 | 2 bytes/elem — halves activation memory |
| gradients (bwd) | FP16 | 2 bytes/elem — halves gradient bandwidth |
| master weights | FP32 | weight updates are tiny; FP16 loses them |
| optimizer state (m, v) | FP32 | momentum accumulators diverge in FP16 |
The memory win is almost entirely in activations. For a 100M-param model with Adam: params + opt state stay ~1.6 GB regardless; a large batch's activations halve.
Discrimination
Sort into buckets
Sort these by dtype, from memory, without looking back at Why mixed precision helps. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Concept
torch.autocast wraps the forward + backward and auto-casts ops to the most efficient dtype — no manual .half() calls needed.
| op category | dtype in autocast | why |
|---|---|---|
| matmul, conv, linear | FP16 | hardware-accelerated on tensor cores |
| batch norm, layer norm | FP32 | statistics accumulation needs precision |
| softmax, log_softmax | FP32 | exponential sum is numerically sensitive |
| loss functions | FP32 | scalar loss must stay accurate for backprop |
Trade off
Comparison matrix
From autocast: letting PyTorch pick the dtype: every row here is a choice with a cost. Fill the why column, then say which row you would actually pick and what you give up for it.
| op category | dtype in autocast | why |
|---|---|---|
| matmul, conv, linear | FP16 | hardware-accelerated on tensor cores |
| batch norm, layer norm | FP32 | statistics accumulation needs precision |
| softmax, log_softmax | FP32 | exponential sum is numerically sensitive |
| loss functions | FP32 | scalar loss must stay accurate for backprop |
Section
Part 2 of 4
Concept
Late in training, gradients for early layers can be very small — easily below FP16's smallest normal (6.10e-05). They underflow to 0 and the layer stops learning.
| gradient value | FP32 stores? | FP16 stores? |
|---|---|---|
| 1.0e-04 | yes 1.000e-04 | yes 1.000e-04 |
| 6.1e-05 | yes 6.100e-05 | yes 6.098e-05 |
| 1.0e-08 | yes 1.000e-08 | NO 0.000e+00 |
The fix: multiply the loss by a large scale S before backward() so all gradients are S× bigger during the backward pass, then divide them back before the optimizer step.
Pattern
Step through it
Step through The underflow problem one row at a time. What is driving the change, and what would the row after the last one be?
Estimation
Predict first
Trace the scale-before / unscale-after pattern. Scale S = 65536 (2^16 = GradScaler's default initial value).
Commit before you compute: what does Loss scaling by hand come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Unscaled gradient = 6.554e-04 / 65536 ≈ 1.00e-08 — correctly recovered
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. Scaling shifts the gradient into FP16's representable range.
Worked example
Trace the scale-before / unscale-after pattern. Scale S = 65536 (2^16 = GradScaler's default initial value).
import torch, torch.nn as nn
S = 65536 # loss scale
g = torch.tensor(1e-8) # tiny gradient, would underflow in FP16
print('g in FP32 :', g.item()) # 1.0e-08
print('g in FP16 :', g.half().item()) # 0.0 -- lost!
print()
scaled_g = g * S
print('g * S in FP32 :', scaled_g.item()) # 6.554e-04
print('g * S in FP16 :', scaled_g.half().item()) # 6.554e-04 -- survived!
print('unscaled (/ S) :', scaled_g.half().item() / S) # ~1e-08 recoveredUnscaled gradient = 6.554e-04 / 65536 ≈ 1.00e-08 — correctly recovered
Why: Scaling shifts the gradient into FP16's representable range. Dividing back after the backward pass restores the true value before the optimizer updates weights.
| quantity | FP32 | FP16 (no scale) | FP16 (scaled, S=65536) |
|---|---|---|---|
| g = 1e-8 | 1.0e-08 | 0.0 (lost) | — |
| g * S = 0.000655 | — | — | 6.554e-04 (safe) |
| unscaled | — | — | 1.00e-08 (recovered) |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
Unscaled gradient = 6.554e-04 / 65536 ≈ 1.00e-08 — correctly recovered
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 scale-before / unscale-after pattern. Scale S = 65536 (2^16 = GradScaler's default initial value).
Concept
torch.amp.GradScaler automates the scale-and-unscale loop. It also detects overflowed gradients (inf/NaN) and skips the optimizer step that iteration.
scaler.step() is a no-op — the update doesn't happenThe scale adapts automatically. You never tune it; just use the loop and let the scaler track the stable range.
Counterexample
Discussion prompt
torch.amp.GradScaler automates the scale-and-unscale loop. It also detects overflowed gradients (inf/NaN) and skips the optimizer step that iteration.
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:
The scale adapts automatically. You never tune it; just use the loop and let the scaler track the stable range.
Missing information
Discussion prompt
The canonical mixed-precision loop adds three things to the standard loop (Lesson 40): autocast, scaler.scale(loss), and scaler.step / scaler.update.
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:
Clean training on CPU — no overflow, so the scale never needs to shrink. On GPU with real FP16 it would adapt within the first few steps.
Worked example
The canonical mixed-precision loop adds three things to the standard loop (Lesson 40): autocast, scaler.scale(loss), and scaler.step / scaler.update.
import torch, torch.nn as nn
torch.manual_seed(42)
X = torch.randn(200, 20); y = (X[:, :10].sum(1) > 0).float().unsqueeze(1)
net = nn.Sequential(nn.Linear(20, 64), nn.ReLU(), nn.Linear(64, 1))
optimizer = torch.optim.Adam(net.parameters(), lr=1e-3)
scaler = torch.amp.GradScaler('cpu')
loss_fn = nn.BCEWithLogitsLoss()
for ep in range(5):
optimizer.zero_grad()
with torch.autocast('cpu'):
pred = net(X)
loss = loss_fn(pred, y)
scaler.scale(loss).backward() # backward on scaled loss
scaler.step(optimizer) # unscale + step (skips if nan/inf)
scaler.update() # adjust scale factor
print(f'ep {ep+1}: loss={loss.item():.4f} scale={int(scaler.get_scale())}')Epoch 1 loss = 0.6901, epoch 5 loss = 0.6791, scale stays at 65536 (no overflow)
Why: Clean training on CPU — no overflow, so the scale never needs to shrink. On GPU with real FP16 it would adapt within the first few steps.
| epoch | loss | scale |
|---|---|---|
| 1 | 0.6901 | 65536 |
| 2 | 0.6874 | 65536 |
| 3 | 0.6845 | 65536 |
| 4 | 0.6819 | 65536 |
| 5 | 0.6791 | 65536 |
Invariant
Step through it
Step through The AMP training loop one row at a time. One of these columns never changes — find it, and say why it cannot.
Anomaly
Predict first
A student writes this, and it looks reasonable:
After scaler.step(optimizer), the model is updated — the loop is done for this step.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: The GradScaler's internal scale factor never changes.
Always call scaler.update() after scaler.step(optimizer) each iteration.
Why: The GradScaler's internal scale factor never changes. After any overflow step, the scale was supposed to halve — but without update() it stays dangerously high, causing repeated inf gradients and skipped optimizer steps.
Trap
After scaler.step(optimizer), the model is updated — the loop is done for this step.
Omit scaler.update() at the end of the loop
Why: The GradScaler's internal scale factor never changes. After any overflow step, the scale was supposed to halve — but without update() it stays dangerously high, causing repeated inf gradients and skipped optimizer steps.
Always call scaler.update() after scaler.step(optimizer) each iteration.
zero_grad -> autocast(forward) -> scaler.scale(loss).backward() -> scaler.step(opt) -> scaler.update()
Why: update() adjusts the scale up (on clean steps) or down (after overflow), keeping gradients in FP16's safe range on every future step.
Break the constraint
Discussion prompt
The rule this trap just fixed:
Always call scaler.update() after scaler.step(optimizer) each iteration.
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:
The GradScaler's internal scale factor never changes. After any overflow step, the scale was supposed to halve — but without update() it stays dangerously high, causing repeated inf gradients and skipped optimizer steps.
Section
Part 3 of 4
Concept
Large batches stabilize the gradient estimate (lower variance per step) and sometimes improve final convergence — but a batch of 1024 may not fit in GPU memory.
Gradient accumulation: run N mini-batches of size B without stepping the optimizer, then step once. The accumulated gradient equals the gradient on a single batch of size N×B.
\[ g_{\text{accum}} = \sum_{i=1}^{N} \nabla_{\theta}\, \frac{1}{N}\mathcal{L}(\theta, \mathcal{B}_i) = \nabla_{\theta}\, \mathcal{L}(\theta, \mathcal{B}_{\text{full}}) \]
Analogy
Discussion prompt
Explain Why accumulate? 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:
Large batches stabilize the gradient estimate (lower variance per step) and sometimes improve final convergence — but a batch of 1024 may not fit in GPU memory.
Estimation
Predict first
128 samples, 2 features, logistic regression. Accumulate 4 mini-batches of 32 vs single pass on 128. Weights must match exactly.
Commit before you compute: what does Accumulation = full batch: the proof come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Key: divide loss by N_ACCUM before backward — else each mini-batch contributes a full loss instead of 1/N of one
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. BCEWithLogitsLoss uses reduction='mean' inside each mini-batch.
Worked example
128 samples, 2 features, logistic regression. Accumulate 4 mini-batches of 32 vs single pass on 128. Weights must match exactly.
import torch, torch.nn as nn
torch.manual_seed(42)
X = torch.randn(128, 2); y = (X[:,0]+X[:,1]>0).float().unsqueeze(1)
loss_fn = nn.BCEWithLogitsLoss()
N_ACCUM, BS = 4, 32
# --- gradient accumulation ---
model = nn.Linear(2, 1, bias=False); nn.init.zeros_(model.weight)
opt = torch.optim.SGD(model.parameters(), lr=0.1)
opt.zero_grad()
for i in range(N_ACCUM):
xb, yb = X[i*BS:(i+1)*BS], y[i*BS:(i+1)*BS]
(loss_fn(model(xb), yb) / N_ACCUM).backward() # divide before backward!
opt.step()
print('accum w:', model.weight.data.round(decimals=5))Key: divide loss by N_ACCUM before backward — else each mini-batch contributes a full loss instead of 1/N of one
Why: BCEWithLogitsLoss uses reduction='mean' inside each mini-batch. We want the mean over the full 128, so each of the 4 mini-batch terms should count 1/4 of its already-averaged loss.
| step | scaled_loss | accum_grad_w0 | accum_grad_w1 |
|---|---|---|---|
| 1 | 0.17329 | -0.03863 | -0.10888 |
| 2 | 0.17329 | -0.12513 | -0.15072 |
| 3 | 0.17329 | -0.23454 | -0.19790 |
| 4 | 0.17329 | -0.30851 | -0.27169 |
Comparison
Comparison matrix
From Accumulation = full batch: the proof: refill the scaled_loss column from what you know. The rest of the table is as it appeared.
| step | scaled_loss | accum_grad_w0 | accum_grad_w1 |
|---|---|---|---|
| 1 | 0.17329 | -0.03863 | -0.10888 |
| 2 | 0.17329 | -0.12513 | -0.15072 |
| 3 | 0.17329 | -0.23454 | -0.19790 |
| 4 | 0.17329 | -0.30851 | -0.27169 |
Estimation
Predict first
Run the single full-batch reference and confirm the weights match to 8 decimal places.
Commit before you compute: what does Verifying the equivalence come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Both methods produce w = [[0.03085, 0.02717]] — difference = 0.0 (bitwise identical)
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. PyTorch accumulates gradients with += into the .grad tensors.
Worked example
Run the single full-batch reference and confirm the weights match to 8 decimal places.
import torch, torch.nn as nn
torch.manual_seed(42)
X = torch.randn(128, 2); y = (X[:,0]+X[:,1]>0).float().unsqueeze(1)
loss_fn = nn.BCEWithLogitsLoss()
# --- single full-batch reference ---
ref = nn.Linear(2, 1, bias=False); nn.init.zeros_(ref.weight)
opt_ref = torch.optim.SGD(ref.parameters(), lr=0.1)
opt_ref.zero_grad()
loss_fn(ref(X), y).backward()
opt_ref.step()
print('single w:', ref.weight.data.round(decimals=5))
# accum w from prev slide: [[0.03085, 0.02717]]
# single w: [[0.03085, 0.02717]] <-- must matchBoth methods produce w = [[0.03085, 0.02717]] — difference = 0.0 (bitwise identical)
Why: PyTorch accumulates gradients with += into the .grad tensors. Dividing each mini-batch loss by N_ACCUM ensures the sum equals the mean over the full batch, matching the single-pass gradient exactly.
| method | w[0] | w[1] | match? |
|---|---|---|---|
| 4-step accumulation | 0.03085 | 0.02717 | |
| single full batch | 0.03085 | 0.02717 | yes (diff = 0.0) |
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
Both methods produce w = [[0.03085, 0.02717]] — difference = 0.0 (bitwise identical)
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:
Run the single full-batch reference and confirm the weights match to 8 decimal places.
Anomaly
Predict first
A student writes this, and it looks reasonable:
In the accumulation loop, backward on the raw loss each mini-batch — grads just add up.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Each mini-batch contributes its own full mean loss, so N iterations accumulate N× the intended gradient.
Divide the loss by N_ACCUM before calling backward each mini-batch.
Why: Each mini-batch contributes its own full mean loss, so N iterations accumulate N× the intended gradient. The effective learning rate is multiplied by N — causing overshooting or divergence.
Trap
In the accumulation loop, backward on the raw loss each mini-batch — grads just add up.
for i in range(N): loss_fn(model(xb), yb).backward() # NO division
Why: Each mini-batch contributes its own full mean loss, so N iterations accumulate N× the intended gradient. The effective learning rate is multiplied by N — causing overshooting or divergence.
Divide the loss by N_ACCUM before calling backward each mini-batch.
(loss_fn(model(xb), yb) / N_ACCUM).backward() # correct
Why: This weights each mini-batch's gradient at 1/N of its per-sample mean, so the accumulated gradient equals the full-batch mean gradient exactly.
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.
torch.autocast wraps the forward + backward and auto-casts ops to the most efficient dtype — no manual .half() calls needed.scaler.step(optimizer), the model is updated — the loop is done for this step.; In the accumulation loop, backward on the raw loss each mini-batch — grads just add up.Section
Part 4 of 4
Concept
A common misconception: AMP halves all training memory. In practice the parameter/optimizer budget is nearly unchanged — it's activation memory that halves.
| tensor | FP32 Adam (bytes/param) | AMP Adam (bytes/param) |
|---|---|---|
| weights | 4 | 2 (fp16) + 4 (fp32 master) = 6 |
| gradients | 4 | 2 (fp16 grads used in bwd) |
| optimizer m, v | 4 + 4 = 8 | 4 + 4 = 8 (always fp32) |
| total params+opt | 16 | ~18 (slightly MORE!) |
The real win: activations (the layer outputs saved for backward) are stored in FP16. For a batch of 64 tokens × 512 seq × 256 hidden: FP32 = 16.8 MB/layer, FP16 = 8.4 MB/layer — 50% per layer.
Discrimination
Sort into buckets
Sort these by FP32 Adam (bytes/param), from memory, without looking back at Where AMP memory is (and isn't) saved. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Constraint
Discussion prompt
Run Mixed precision + gradient accumulation recipe with this step confiscated:
Step + update with scaler.step(opt); scaler.update() — unscales grads, skips overflow steps, adapts scale
Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.
Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.
Answer:
with torch.autocast('cuda'): — ops choose FP16/FP32 automaticallyscaler.scale(loss).backward() — prevents FP16 gradient underflowscaler.step(opt); scaler.update() — unscales grads, skips overflow steps, adapts scaleloss / N_ACCUM before backward; call opt.step() and opt.zero_grad() only after all N mini-batchesscaler.scale(loss/N).backward() each mini-batch; only scaler.step and scaler.update after the N-thPattern
with torch.autocast('cuda'): — ops choose FP16/FP32 automaticallyscaler.scale(loss).backward() — prevents FP16 gradient underflowscaler.step(opt); scaler.update() — unscales grads, skips overflow steps, adapts scaleloss / N_ACCUM before backward; call opt.step() and opt.zero_grad() only after all N mini-batchesscaler.scale(loss/N).backward() each mini-batch; only scaler.step and scaler.update after the N-thEdge cases
Discussion prompt
Mixed precision + gradient accumulation 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:
with torch.autocast('cuda'): — ops choose FP16/FP32 automaticallyscaler.scale(loss).backward() — prevents FP16 gradient underflowscaler.step(opt); scaler.update() — unscales grads, skips overflow steps, adapts scaleloss / N_ACCUM before backward; call opt.step() and opt.zero_grad() only after all N mini-batchesscaler.scale(loss/N).backward() each mini-batch; only scaler.step and scaler.update after the N-thElimination
Eliminate the wrong options
A gradient of 1e-08 is cast to FP16. What is the stored value?
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: FP16's smallest representable normal is 6.10e-05; 1e-08 is below even the subnormal range (~6e-08). torch.tensor(1e-8).half().item() == 0.0. Loss scaling multiplies it by ~65536 before the cast to keep it alive.
Check
Think about the FP16 table from earlier.
Check your understanding
A gradient of 1e-08 is cast to FP16. What is the stored value?
Answer: A
Why: FP16's smallest representable normal is 6.10e-05; 1e-08 is below even the subnormal range (~6e-08). torch.tensor(1e-8).half().item() == 0.0. Loss scaling multiplies it by ~65536 before the cast to keep it alive.
Prediction
Predict first
Which is the correct order of one AMP training step?
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: zero_grad → autocast(forward+loss) → scaler.scale(loss).backward() → scaler.step(opt) → scaler.update()
Why: zero_grad first (clear old grads); forward + loss inside autocast; backward on the scaled loss; scaler.step unscales grads and calls opt.step if no overflow; scaler.update adjusts the scale factor.
Check
Reconstruct the AMP training loop order.
Check your understanding
Which is the correct order of one AMP training step?
Answer: A
Why: zero_grad first (clear old grads); forward + loss inside autocast; backward on the scaled loss; scaler.step unscales grads and calls opt.step if no overflow; scaler.update adjusts the scale factor.
Elimination
Eliminate the wrong options
You want to accumulate gradients over 4 mini-batches to simulate batch-128 training (32 per mini-batch). Which backward call is correct?
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: With reduction='mean', each mini-batch BCELoss already averages over its 32 samples. To get the full-128-sample mean gradient, each mini-batch's contribution must be 1/4 of its per-sample mean. Dividing before backward achieves this exactly — verified: weight diff vs single full batch = 0.0.
Check
A concrete scenario: 4 mini-batches, BCEWithLogitsLoss(reduction='mean').
Check your understanding
You want to accumulate gradients over 4 mini-batches to simulate batch-128 training (32 per mini-batch). Which backward call is correct?
Answer: A
Why: With reduction='mean', each mini-batch BCELoss already averages over its 32 samples. To get the full-128-sample mean gradient, each mini-batch's contribution must be 1/4 of its per-sample mean. Dividing before backward achieves this exactly — verified: weight diff vs single full batch = 0.0.
Section
Project
Concept
Build a training loop that combines mixed precision (autocast + GradScaler) with gradient accumulation (N=4 mini-batches). Verify the final weights match a plain FP32 full-batch reference.
| # | requirement | key call |
|---|---|---|
| 1 | GradScaler + autocast loop | torch.amp.GradScaler; torch.autocast |
| 2 | gradient accumulation (N=4) | divide loss by N before backward |
| 3 | verify weight match vs reference | torch.allclose on final weights |
Build rules: use a small synthetic dataset (128 × 10 features, binary labels). Do NOT call opt.step() or opt.zero_grad() inside the accumulation loop — only after all N mini-batches.
Counterexample
Discussion prompt
Build a training loop that combines mixed precision (autocast + GradScaler) with gradient accumulation (N=4 mini-batches). Verify the final weights match a plain FP32 full-batch reference.
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: use a small synthetic dataset (128 × 10 features, binary labels). Do NOT call opt.step() or opt.zero_grad() inside the accumulation loop — only after all N mini-batches.
Worked example
Your turn: write the 5-line AMP loop skeleton. Predict: does the loss scale change in the first 5 epochs on this small clean problem?
Hint: scaler = torch.amp.GradScaler('cpu'); wrap forward+loss in with torch.autocast('cpu'):; call scaler.scale(loss).backward() then scaler.step(opt).
import torch, torch.nn as nn
torch.manual_seed(42)
X = torch.randn(128, 10); y = (X[:,:5].sum(1)>0).float().unsqueeze(1)
net = nn.Sequential(nn.Linear(10,32), nn.ReLU(), nn.Linear(32,1))
opt = torch.optim.Adam(net.parameters(), lr=1e-3)
scaler = torch.amp.GradScaler('cpu')
for ep in range(5):
opt.zero_grad()
with torch.autocast('cpu'):
loss = nn.BCEWithLogitsLoss()(net(X), y)
scaler.scale(loss).backward()
scaler.step(opt); scaler.update()
print(f'ep {ep+1} loss={loss.item():.4f} scale={int(scaler.get_scale())}')| epoch | expected loss range | expected scale |
|---|---|---|
| 1 | ~0.69 (near random) | 65536 (no overflow) |
| 2-5 | slowly decreasing | 65536 (stays stable) |
Trade off
Comparison matrix
From Milestone 1 — AMP loop (no accumulation yet): every row here is a choice with a cost. Fill the expected scale column, then say which row you would actually pick and what you give up for it.
| epoch | expected loss range | expected scale |
|---|---|---|
| 1 | ~0.69 (near random) | 65536 (no overflow) |
| 2-5 | slowly decreasing | 65536 (stays stable) |
Worked example
Your turn: extend the loop to accumulate over N_ACCUM=4 mini-batches of 32 each. Only step and zero_grad after all 4 mini-batches.
Hint: opt.zero_grad() before the inner loop; inside divide loss / N_ACCUM before scaler.scale(...).backward(); call scaler.step and scaler.update once after the inner loop.
import torch, torch.nn as nn
torch.manual_seed(42)
X = torch.randn(128, 10); y = (X[:,:5].sum(1)>0).float().unsqueeze(1)
net = nn.Sequential(nn.Linear(10,32), nn.ReLU(), nn.Linear(32,1))
opt = torch.optim.Adam(net.parameters(), lr=1e-3)
scaler = torch.amp.GradScaler('cpu')
N_ACCUM, BS = 4, 32
opt.zero_grad()
for i in range(N_ACCUM):
xb, yb = X[i*BS:(i+1)*BS], y[i*BS:(i+1)*BS]
with torch.autocast('cpu'):
loss = nn.BCEWithLogitsLoss()(net(xb), yb) / N_ACCUM
scaler.scale(loss).backward()
scaler.step(opt); scaler.update()
for p in list(net.parameters())[:2]:
print(p.data.flatten()[:3].round(decimals=4))| check | expected |
|---|---|
| opt.step() called | once, after i=0..3 loop |
| opt.zero_grad() called | once, before the inner loop |
| loss division | / N_ACCUM inside the loop |
Comparison
Comparison matrix
From Milestone 2 — add gradient accumulation: refill the expected column from what you know. The rest of the table is as it appeared.
| check | expected |
|---|---|
| opt.step() called | once, after i=0..3 loop |
| opt.zero_grad() called | once, before the inner loop |
| loss division | / N_ACCUM inside the loop |
Concept
Out loud, slides closed: (1) name the 5 steps of the AMP training loop, (2) explain why scaler.update() is non-negotiable, and (3) derive why you divide loss by N_ACCUM and not gradient by N_ACCUM after the loop.
Stretch (homework): run the Lesson 40 linear regression in AMP mode. Compare training loss curves: FP32 vs AMP, 200 epochs. Does the GradScaler ever trigger a backoff? When does large-batch training (from gradient accumulation) hurt generalization? Research sharp vs flat minima (Keskar et al. 2017).
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — FP16 & FP32: what they can hold · Loss scaling: preventing underflow · Gradient accumulation: simulating large batches · Memory budget & the full picture · Your turn: AMP + accumulation. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
autocast → scaler.scale(loss).backward() → scaler.step → scaler.update()| idea | the one thing to remember |
|---|---|
| FP16 underflow | gradients < 6e-05 become 0 without scaling |
| loss scaling | scale up before bwd, unscale inside scaler.step |
| GradScaler | scaler.update() adapts the scale each step — never skip it |
| grad accumulation | divide loss by N_ACCUM before backward each mini-batch |
| AMP memory win | activations halve; optimizer state stays FP32 |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.