Mixed precision training is the default way large models are trained today, and it is also one of the easiest things to get subtly wrong. The idea is simple: run the expensive matrix multiplications in a 16-bit format that the GPU's tensor cores execute many times faster than FP32, while keeping the numerically fragile parts, the weights, the optimizer state and a handful of reductions, in 32-bit. Done right, a training step gets markedly faster, activation memory roughly halves, and the loss curve is indistinguishable from an FP32 run. Done wrong, the loss plateaus, spikes to NaN at step 40,000, or the run is no faster at all.
The number formats themselves, why FP16 underflows, why BF16 has range but little precision and how loss scaling rescues small gradients, are covered from first principles in Mixed precision training: the numerics. This article covers the GPU engineering on top: what PyTorch's automatic mixed precision (AMP) does to each op, choosing a dtype for your hardware, the real memory budget, sharded training, debugging NaNs and proving the speedup.
What the GPU gains and what it risks
Three separate wins come from lower precision on a GPU. The first is compute: tensor cores execute matrix multiply-accumulate on 16-bit inputs at a much higher rate than the FP32 path, and the accumulator inside the instruction stays FP32, so a long dot product does not lose precision as it grows. How the instruction is structured is covered in tensor core architecture. The second is bandwidth: every tensor read from HBM is half the bytes, which matters for the many memory-bound ops such as elementwise activations and normalisation. The third is capacity: activations saved for the backward pass are stored in the lower precision, so a larger batch or a longer sequence fits.
The risks are equally specific. FP16 overflows above 65,504 and flushes gradients below roughly 6e-8 to zero. BF16 has FP32's range but only seven mantissa bits, so a small update added to a large BF16 weight can round away. The cure for both: keep a 32-bit copy of anything that accumulates over time, and 32-bit math for rounding-sensitive ops.
The data flow of one AMP step
Follow the diagram from the top left. The parameters are FP32 and they are the only persistent copy of the weights; this is the FP32 master copy that older recipes kept by hand. Inside the torch.autocast region, each op consults a per-op policy. Ops on the lower-precision list, the matrix multiplications, linear layers, convolutions and batched GEMMs, cast their inputs to the 16-bit type and run on tensor cores. Ops on the FP32 list, such as softmax, log-softmax, layer and group normalisation, exponentials, and the loss functions, cast up and run in FP32. Everything else runs in whatever type its inputs already have, promoting to the widest input when types are mixed.
The backward pass runs outside the autocast block but mirrors the forward: each op's gradient runs in the dtype its forward used, and because the leaf parameters are FP32, the gradients accumulated into .grad are FP32. The optimizer reads FP32 gradients and updates FP32 weights and moments. FP16 adds two boxes, scaling the loss before backward and unscaling plus an inf check before the step; BF16 needs neither.
Choosing FP16, BF16 or TF32
| Format | Exponent / mantissa bits | Tensor cores from | Needs loss scaling | Use it when |
|---|---|---|---|---|
| FP16 | 5 / 10 | Volta (V100) | Yes, GradScaler | Pre-Ampere GPUs, or inference-matched training |
| BF16 | 8 / 7 | Ampere (A100) | No | The default for training on A100, H100 and newer |
| TF32 | 8 / 10 (internal) | Ampere (A100) | No | FP32 code you cannot change; a free speedup for GEMMs |
| FP8 (E4M3 / E5M2) | 4/3 or 5/2 | Hopper (H100) | Per-tensor scaling | Large models with a library such as Transformer Engine |
On any GPU that supports it, BF16 is the right starting point. It removes an entire class of failure, gradient underflow and overflow, and with it the GradScaler machinery. Its lower mantissa precision is harmless for activations and GEMM inputs because the accumulators are FP32 and the weights are updated in FP32. Compute capability 8.0 or higher, from torch.cuda.get_device_capability(), means BF16 tensor cores are present.
FP16 still has a place: on Volta and Turing it is the only tensor-core training format, and it has three more mantissa bits than BF16. TF32 is not a storage format at all: it is a mode in which FP32 matmuls are executed on tensor cores with a reduced mantissa. Enable it with torch.set_float32_matmul_precision("high"); it speeds up any FP32 GEMM that still runs outside autocast. FP8 needs scaling per tensor rather than per loss and is best left to a library that manages the scale factors for you.
A reference training loop
The loop below is the pattern to copy. It selects BF16 when the device supports it and falls back to FP16 with a scaler, uses gradient accumulation, clips gradients correctly and checkpoints the scaler. With enabled=False the scaler's methods become pass-throughs, so the same loop serves both dtypes.
import torch
from torch import nn
device = "cuda"
use_bf16 = torch.cuda.get_device_capability()[0] >= 8 # Ampere+: BF16 tensor cores
amp_dtype = torch.bfloat16 if use_bf16 else torch.float16
model = MyModel().to(device) # parameters stay float32
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, fused=True)
scaler = torch.amp.GradScaler("cuda", enabled=not use_bf16)
torch.set_float32_matmul_precision("high") # TF32 for FP32 matmuls outside autocast
accum = 4 # micro-batches per optimizer step
for step, (x, y) in enumerate(loader):
x, y = x.to(device, non_blocking=True), y.to(device, non_blocking=True)
with torch.autocast(device_type="cuda", dtype=amp_dtype):
logits = model(x) # GEMMs in amp_dtype
loss = nn.functional.cross_entropy(logits, y) / accum # autocast runs it in FP32
scaler.scale(loss).backward() # backward OUTSIDE autocast
if (step + 1) % accum == 0:
scaler.unscale_(opt) # real gradient values, once per step
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(opt) # skipped if any grad is inf/NaN
scaler.update() # halves or grows the scale
opt.zero_grad(set_to_none=True)
# checkpoint the scaler with everything else, or a resume restarts at 65536
ckpt = {"model": model.state_dict(), "opt": opt.state_dict(),
"scaler": scaler.state_dict(), "step": step}Four details in it are where most hand-written loops go wrong. Only the forward pass and the loss are inside autocast; wrapping backward is unnecessary and not recommended. scaler.unscale_(opt) must be called before clip_grad_norm_, otherwise you clip scaled gradients against an unscaled threshold and the clip is either a no-op or a disaster. Unscale happens once per optimizer step, not once per micro-batch. And the model is never converted with .half(): autocast needs FP32 parameters to produce FP32 gradients, and a model that is already FP16 gives you pure FP16 training with no master copy.
GradScaler's defaults are sensible: an initial scale of 2 to the 16th, halved whenever a step finds inf or NaN, and doubled after 2,000 consecutive clean steps. A few skipped steps early in training are normal while the scale settles.
What autocast actually does to your code
Autocast is a dispatcher-level policy, not a model transformation, and that has consequences. It applies to ops executed inside the context on the matching device type, so a CPU op inside a CUDA autocast region is untouched. It caches the low-precision copy of each FP32 parameter for the duration of the region, so a weight used twice in a forward pass is cast once, and the cache is dropped when the region exits. This is also why the context should wrap the whole forward pass once, rather than being entered and exited per layer.
Custom code needs care. A hand-written CUDA or Triton kernel called inside autocast receives whatever dtype the preceding op produced, which may not be what the kernel expects. Decorate custom autograd functions with torch.amp.custom_fwd(device_type="cuda", cast_inputs=torch.float32) to force a dtype, or with custom_fwd and custom_bwd without casting to make the backward run under the same autocast state as the forward. For a numerically sensitive block inside an autocast region, such as a custom attention score computation, open a nested torch.autocast(device_type="cuda", enabled=False) and cast its inputs to FP32 explicitly.
Two quieter effects are worth knowing. Tensors you create yourself inside the region, such as masks and position encodings, keep their own dtype, and adding an FP32 tensor to a BF16 activation promotes the result to FP32, silently doubling memory for everything downstream.
The memory budget, worked through
Take a 1.3 billion parameter model trained with AdamW. A common misconception is that mixed precision halves the whole footprint. It does not, because the persistent state stays FP32: 4 bytes for each weight, 4 for each gradient and 8 for Adam's two moments, 16 bytes per parameter in total, about 20.8 GB. Autocast adds a transient BF16 copy of the weights during forward and backward, another 2 bytes per parameter, about 2.6 GB at peak.
What mixed precision shrinks is the activation memory, which for transformers at realistic batch sizes and sequence lengths is often larger than all of the above. Activations saved for backward are stored in BF16, 2 bytes instead of 4 per element, so a batch that needed 40 GB of activations in FP32 needs about 20 GB. That is where the headroom for a bigger batch or a longer context comes from. Where the rest of a step's memory goes is itemised in anatomy of a GPU training step.
Pure BF16 training, storing the parameters themselves in BF16, saves more but reintroduces update swamping late in training; if you need that saving, use an optimizer with FP32 master weights or compensated summation rather than casting the model.
Mixed precision in sharded training
Under plain data parallelism nothing changes: each replica runs the loop above, and the gradients being all-reduced are FP32 because the parameters are. Sharded training is different, because the framework owns the parameter storage and moves parameters and gradients across the network. PyTorch FSDP2 therefore takes an explicit policy:
import torch
from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy
mp = MixedPrecisionPolicy(
param_dtype=torch.bfloat16, # all-gathered weights used in forward/backward
reduce_dtype=torch.float32, # gradient reduce-scatter stays in full precision
)
for block in model.layers: # shard each transformer block, then the root
fully_shard(block, mp_policy=mp)
fully_shard(model, mp_policy=mp)
# The sharded parameters the optimizer sees remain FP32: the FP32 master copy is
# distributed across ranks, and only the gathered working copy is BF16.param_dtype is the dtype of the all-gathered working copy that forward and backward use, so the gather traffic halves. reduce_dtype is the dtype of the gradient reduce-scatter. Reducing in BF16 halves that traffic too, but summing many ranks' gradients in a 7-bit mantissa loses precision as the world size grows; keep it FP32 unless communication is your measured bottleneck. The sharded parameters that the optimizer updates stay FP32, so the master copy survives, split across ranks. The rest of the FSDP2 API is covered in PyTorch FSDP in depth, and DeepSpeed's equivalent settings, including where each ZeRO stage puts the FP32 state, in the ZeRO optimizer.
Failure modes and how to debug them
- Scale keeps falling. With FP16, a scale that halves repeatedly and never recovers means something overflows on nearly every step, often an attention logit or an unnormalised residual stream. Log
scaler.get_scale()every step; a collapse below 1 is a bug, not bad luck. - NaN appears suddenly. Find the first module that produces a non-finite value rather than the step where the loss shows it. The hooks below do that;
torch.autograd.detect_anomaly()does the same for backward but is slow, so enable it only while reproducing. - Loss plateaus early. Usually weights stored in 16-bit, from a stray
.half(),model.to(torch.bfloat16)or a checkpoint loaded in the wrong dtype. Assertnext(model.parameters()).dtype == torch.float32at startup. - No speedup. Tensor cores need suitably aligned shapes; NVIDIA's guidance is to keep GEMM dimensions such as hidden size, vocabulary and batch times sequence at multiples of 8 for 16-bit types. Small models may also be bound by the data loader or kernel-launch overhead, which precision does not help.
- Resume changes behaviour. A checkpoint without the scaler's state restarts at the initial scale, producing a burst of skipped steps. Save and load
scaler.state_dict().
def find_first_nonfinite(model):
"""Forward hooks that name the first module producing inf/NaN."""
def hook(mod, inp, out):
t = out[0] if isinstance(out, tuple) else out
if torch.is_tensor(t) and not torch.isfinite(t).all():
raise RuntimeError(f"non-finite output from {mod.__class__.__name__} "
f"(dtype {t.dtype}, max {t.float().abs().max().item():.3e})")
return [m.register_forward_hook(hook) for m in model.modules()]
Measuring the win
Measure; do not assume. Time steady-state steps with CUDA events, because kernel launches are asynchronous, and record peak memory.
def time_steps(step_fn, n=50, warmup=10):
for _ in range(warmup):
step_fn()
torch.cuda.synchronize()
start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(n):
step_fn()
end.record()
torch.cuda.synchronize()
return start.elapsed_time(end) / n # milliseconds per step
# run once with autocast disabled and once enabled; also record
# torch.cuda.max_memory_allocated() after torch.cuda.reset_peak_memory_stats()Compare three configurations on the same batch: FP32 with TF32 disabled, FP32 with TF32 enabled, and autocast in BF16 or FP16. Compare loss curves over a few thousand steps too, not just step time. Finally, profile one step with torch.profiler and check that the heavy GEMM kernels are tensor-core kernels; if FP32 GEMMs remain, some op is escaping autocast.
What to do next
- Check your GPU's compute capability: use BF16 autocast on 8.0 and newer, FP16 with GradScaler below that.
- Wrap only forward and loss in
torch.autocast(device_type="cuda", dtype=...); keep backward and the optimizer step outside. - Leave parameters FP32; add a startup assertion that they are.
- With FP16, call
scaler.unscale_before clipping, log the scale every step and checkpoint the scaler state. - Enable TF32 for any FP32 matmuls that remain.
- Pad hidden, vocabulary and sequence-batch dimensions to multiples of 8.
- In FSDP2, set
param_dtype=torch.bfloat16and keepreduce_dtype=torch.float32until you have measured communication as the bottleneck. - Benchmark step time, peak memory and a few thousand steps of loss against an FP32 baseline before you declare victory.