scaler = torch.amp.GradScaler('cuda')
opt.zero_grad()
for step, batch in enumerate(loader):
with torch.autocast('cuda'):
loss = model(batch) / accum_steps # divide, or gradients are 4x too big
scaler.scale(loss).backward() # grads accumulate across iterations
if (step + 1) % accum_steps == 0:
scaler.unscale_(opt)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(opt)
scaler.update()
opt.zero_grad(set_to_none=True)
The graded details: (1) divide the loss by accum_steps — the dreaded trap where the per-step LR effectively multiplies by accum_steps; (2) zero_grad only after stepping, since .backward() accumulates by design; (3) scaler.unscale_ before clipping so you clip true gradient norms — otherwise clip happens on scaled grads and you get tiny clipped values; (4) don't call loss.backward() outside autocast on a mixed-precision model; (5) don't recompute the warmup target after each zero_grad.
Explain why GradScaler exists: fp16 has a tiny exponent range, so small gradients underflow to zero; scaling the loss up (and gradients back down before the step) preserves them. bf16 has fp32's exponent range, so on A100+/H100 you use bf16 and drop the scaler entirely.
Follow-ups: Is accumulation exactly equivalent to a bigger batch? (Almost — BatchNorm statistics differ; LayerNorm models like transformers are fine.) Where does the memory actually go? (Params + grads + Adam states ≈ 16 bytes/param; activations scale with batch — hence gradient checkpointing.) When accumulation isn't enough: gradient checkpointing, LoRA/QLoRA, ZeRO/FSDP sharding — ordered by invasiveness.