How do torch.amp.autocast and GradScaler fit into a PyTorch training loop?
answer
- two halves: casting and scaling
- wrap forward and loss only
- 16-bit range is the problem, not precision
- the scaler may skip an update
- bfloat16 makes half of it unnecessary
basics
~20 sRun only the forward pass and the loss under torch.amp.autocast, which executes eligible operations in a 16-bit type while the weights stay 32-bit. With float16, route the backward pass through torch.amp.GradScaler so small gradients do not underflow to zero.
solid answer
~50 sMixed precision has two halves. `torch.amp.autocast(device_type="cuda", dtype=torch.float16)` wraps **only** the forward pass and the loss computation; inside it, PyTorch runs matmuls and convolutions in the low-precision type while keeping reductions and normalizations in float32, and the master weights remain float32 throughout. Backward is left outside the context — it reuses the dtypes autocast chose on the way forward. The second half is loss scaling, needed for float16 only. `torch.amp.GradScaler("cuda")` multiplies the loss by a large factor so that tiny gradients land inside float16's representable range: `scaler.scale(loss).backward()`, then `scaler.step(optimizer)` which unscales the gradients, checks them for inf/NaN and **skips the update** if any are found, then `scaler.update()` which raises or lowers the scale factor for next time. With bfloat16 the scaler is unnecessary — bf16 has float32's exponent range — so people typically construct it with `enabled=False` or skip it. Save `scaler.state_dict()` in checkpoints.
code
python · 20 linesimport torch
from torch import nn
device = "cuda" if torch.cuda.is_available() else "cpu"
use_fp16 = device == "cuda"
model = nn.Linear(16, 4).to(device)
opt = torch.optim.AdamW(model.parameters(), lr=1e-3)
scaler = torch.amp.GradScaler(device, enabled=use_fp16)
x = torch.randn(32, 16, device=device)
y = torch.randint(0, 4, (32,), device=device)
opt.zero_grad(set_to_none=True)
with torch.amp.autocast(device_type=device, dtype=torch.float16, enabled=use_fp16):
loss = nn.functional.cross_entropy(model(x), y)
scaler.scale(loss).backward()
scaler.step(opt)
scaler.update()
print(loss.item(), scaler.get_scale())go deeper
Know that mixed precision exists to make training faster and lighter on GPU memory, and that it is turned on by wrapping the forward pass in a context manager rather than by casting the model yourself.
Explain the placement rules and the four scaler calls in order, and why float16 needs a loss scale while bfloat16 does not. Be able to write the loop from memory.
Talk through skipped steps, scale collapse, checkpointing the scaler state, and the debugging ladder — disable the scaler, then try bfloat16 — when losses go NaN under AMP.
Decide the precision policy for a fleet: which dtype per hardware generation, what accuracy regression gate is required before turning it on, and whether the memory saved is spent on larger batches or longer context.
## What autocast does `torch.amp.autocast` is a context manager that intercepts operator dispatch. Inside it, operations on a per-op allowlist — matrix multiplies, convolutions, linear layers, the heavy compute — are cast to the chosen 16-bit dtype, while operations that are sensitive to precision loss — softmax, layer norm, most reductions, loss functions — stay in float32. You do not cast the model or the inputs yourself; you wrap the region and let the dispatcher decide per operation. The parameters themselves stay in float32. This is the "mixed" in mixed precision: a float32 master copy of the weights, 16-bit arithmetic for the expensive kernels, float32 accumulation inside those kernels on hardware with tensor cores. The payoff is roughly halved activation memory and substantially faster matmuls on modern GPUs. Two placement rules follow directly: - **Only the forward pass and the loss go inside the context.** Backward is deliberately outside; autograd records the dtype each op used going forward and reuses it going backward, so wrapping backward gains nothing and can confuse the scaling logic. - **The optimizer step is outside too**, because it must operate on float32 master weights. ## Why float16 needs a loss scale float16 has 5 exponent bits, giving a smallest normal magnitude around 6e-5. Gradients in a deep network routinely sit below that, and anything smaller flushes to zero — the update simply disappears, and the model trains to a visibly worse result with no error anywhere. Loss scaling fixes it by exploiting linearity: multiply the loss by S before `backward()`, and every gradient in the graph comes out multiplied by S too, shifted up into the representable range. Divide by S before the optimizer sees them and you are back where you started, minus the underflow. bfloat16 has 8 exponent bits — the same range as float32, traded against fewer mantissa bits — so gradients do not underflow and no scaling is needed. That is the main practical reason bf16 is the default choice on hardware that supports it. ## The GradScaler protocol The four calls, in order: 1. `scaler.scale(loss).backward()` — multiplies the loss by the current scale factor, then runs the normal backward. `.grad` now holds scaled gradients. 2. `scaler.unscale_(optimizer)` — *optional*, and only needed if you want to touch the gradients yourself, most commonly to clip them. It divides the gradients in place by the scale factor. 3. `scaler.step(optimizer)` — unscales the gradients if step 2 did not, inspects them for inf or NaN, and calls `optimizer.step()` **only if they are all finite**. If any are not, the step is skipped entirely. 4. `scaler.update()` — adjusts the scale for the next iteration: multiply it up after a run of successful steps, cut it back sharply after a skipped one. That skipping behaviour is the part candidates most often miss. Early in training you should expect a handful of skipped steps while the scaler finds a workable factor — that is normal, not a bug. But it means the number of optimizer steps is not exactly the number of iterations, which slightly desynchronizes anything counting steps, and it means a persistently skipped step (`scaler.get_scale()` collapsing toward 1) is a signal of genuine numerical trouble in the model, not a scaling problem. ## What must stay outside - Backward and the optimizer step, as above. - Anything that manually inspects or edits `.grad` must come after `unscale_`, or it operates on scaled values. - Custom operations that are numerically fragile can be forced back to float32 by wrapping them in `autocast(device_type=..., enabled=False)` and casting the inputs explicitly. - Do not assume tensors coming out of the context are float32 — a value produced inside autocast may be float16, and combining it later with a float32 tensor in an in-place operation can raise a dtype error. ## Checkpoints and diagnostics `GradScaler` has its own `state_dict()`/`load_state_dict()` holding the current scale and growth tracker. Include it in checkpoints alongside the model, optimizer and scheduler, otherwise a resumed run rediscovers the scale from the initial value and may skip a few steps at the seam. When debugging suspected AMP problems, the standard sequence is: disable the scaler and see whether the NaNs persist (if they do, the model is unstable regardless of precision); try bfloat16 instead of float16 (if that fixes it, it was range, not precision); and inspect `scaler.get_scale()` over time to see whether steps are being skipped repeatedly. ## The namespace to use Write `torch.amp.autocast(device_type="cuda", ...)` and `torch.amp.GradScaler("cuda")`. The older device-specific spellings under `torch.cuda.amp` still exist but are the deprecated form; the `torch.amp` namespace is the current one and takes the device type as an argument, which is what makes the same code run on CPU or other backends.
- What does scaler.update() do after a step where the gradients contained inf?It cuts the scale factor by the backoff factor — halving it, by default — so the next iteration scales less aggressively. After a configurable run of consecutive successful steps it instead multiplies the scale up, probing for the largest factor that stays finite. That step itself was already skipped by scaler.step(), so no update reached the weights.
- Why is the backward pass left outside the autocast context?Autograd records the dtype each operation used during the forward pass and runs the corresponding backward in the matching precision automatically. Wrapping backward in autocast adds nothing and is explicitly discouraged in PyTorch's documentation, because the region is meant to bracket only the ops whose dispatch you want redirected.
- Your float16 AMP run produces NaN losses. How do you tell whether AMP is the cause?Rerun with the scaler and autocast disabled: if NaNs persist, the model or data is unstable independently of precision. If they vanish, switch to bfloat16 — same speedup, float32 exponent range — and if that also fixes it the problem was float16 range. Watching scaler.get_scale() collapse toward 1 across steps is the other tell.
saying these in an interview costs you the question
- Wrapping backward() and optimizer.step() inside autocast
- Believing autocast converts the model weights to float16
- Assuming GradScaler is needed for bfloat16 training
- Not knowing the scaler can skip an optimizer step entirely
- Clipping gradients before calling scaler.unscale_()