Autocast can reduce compute and activation cost, but scaled gradients must be unscaled before clipping and the scale must change only at an optimizer-step boundary.
Mixed precision, loss scaling and gradient clipping
Keep a numerical baseline
Train a small fixed batch in full precision before enabling autocast. Record finite loss, gradient norm, predictions and update count. Mixed precision chooses lower precision for eligible operations; it does not mean every tensor or reduction is half precision. Depending on hardware, bfloat16 may need no gradient scaling while float16 often benefits from it. A speedup is hardware- and shape-dependent, so compare measured throughput and accuracy rather than assuming a universal gain. The logits-and-loss contract stays the same.
Order the update deliberately
Run forward and loss under autocast, scale the loss, then run backward. For float16, wait until all microbatches for the current effective batch have contributed before unscaling. Clip the unscaled gradients and call the scaler-controlled optimizer step; it can skip an unsafe update when gradients are nonfinite. Update the scale once at that boundary. Clearing gradients and advancing a scheduler tied to real optimizer updates require an explicit rule for a skipped step.
Understand overflow evidence
A falling scale or repeated skipped updates can mean unstable activations, an excessive learning rate or a model that does not fit float16 range. Log loss, unscaled norm, scale and skipped-update count without dumping sensitive training examples. Switching precision or reducing learning rate may be warranted, but first locate the first nonfinite tensor. Clipping after scaling changes the intended threshold by the scale factor and can suppress almost every useful update.
Account for resumed state
A checkpoint that resumes mixed-precision training should carry model weights, optimizer, scheduler, scaler, update number and data position as one coherent generation. Saving only weights changes future steps even if the first forward output matches. If an update was skipped, the recorded optimizer step must not advance as though parameters changed. Checkpoint recovery covers the publication boundary for a complete state bundle.
Test on the deployment device
Autocast support, kernels, memory behavior and precision vary by accelerator. Profile the same batch shape and sequence length that production uses, including a short final batch and a long tail example. Compare full precision and mixed precision on a frozen validation set, not on a repeatedly tuned test set. A tiny numerical difference is expected; a label flip near a decision threshold needs a product-level tolerance and measured frequency.
Implementation
import torch
from torch import nn
device = torch.device("cuda")
if not torch.cuda.is_available():
raise RuntimeError("This float16 example requires a CUDA device")
quality_model = nn.Linear(3, 2).to(device)
optimizer = torch.optim.AdamW(quality_model.parameters(), lr=0.00047)
scaler = torch.amp.GradScaler("cuda")
features = torch.tensor([[0.7, 0.2, 0.1], [0.1, 0.6, 0.3]], device=device)
labels = torch.tensor([0, 1], device=device)
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type="cuda", dtype=torch.float16):
logits = quality_model(features)
loss = nn.functional.cross_entropy(logits, labels)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
gradient_norm = nn.utils.clip_grad_norm_(quality_model.parameters(), 1.7)
assert torch.isfinite(gradient_norm)
scaler.step(optimizer)
scaler.update()Performance and operating cost
Autocast can reduce eligible activation sizes and accelerate supported kernels, while scaler bookkeeping and some full-precision operations remain. Loss scaling does not repair a truly divergent model. Unscale and norm calculation inspect all gradients, costing O(P) for P parameters; clipping does not lower activation memory. Measure memory and examples per second on the actual device, including skipped updates.
Common Mistakes
- Do not clip scaled gradients as though they were ordinary gradients.
- Do not update the scale in the middle of one effective batch.
- Do not treat every skipped step as a successful parameter update in a scheduler or checkpoint.
Read next
- Gradient accumulation and effective batch accounting
- Checkpoint recovery: save optimizer state and the run boundary
- Training and validation modes: measure the model you will serve
- Logits, cross-entropy and gradients: align the training calculation
- Project: train a receipt model under a memory budget
Continue the workflow: Project: localize receipt defects with a residual CNN.
Continue the workflow: Variational latent sampling, KL accounting and collapse.
Continue the workflow: Distributed checkpoint manifests and coherent resume.
Continue the workflow: AdamW decay groups and update-indexed schedules.
Continue the workflow: Training memory profiling and allocator boundaries.
