Modern GPUs contain specialised units (Tensor Cores and their equivalents) that perform matrix multiplication far faster in 16-bit or lower precision than in 32-bit. Mixed-precision training exploits this: most computation happens in 16-bit, while numerically sensitive parts stay in 32-bit. It typically roughly halves activation memory and speeds up training substantially, usually with no loss in accuracy. It is standard practice for training any sizeable model.
Floating-point formats#
| Format | Exponent bits | Mantissa bits | Range (approx.) | Precision (decimal digits) |
|---|---|---|---|---|
| FP32 | 8 | 23 | $10^{-38}$ to $3 \times 10^{38}$ | ~7 |
| FP16 | 5 | 10 | $6 \times 10^{-5}$ (normal) to 65,504 | ~3 |
| BF16 | 8 | 7 | same as FP32 | ~2โ3 |
| FP8 (E4M3 / E5M2) | 4 / 5 | 3 / 2 | narrow | ~1 |
- FP16 has decent precision but a narrow range: small gradients underflow to zero, large activations overflow to infinity.
- BF16 ("brain float") keeps FP32's exponent โ same range โ with less precision. Overflow and underflow are rare, which makes it the preferred format for training on hardware that supports it.
- FP8 is used on the newest accelerators for further speed-ups, with careful per-tensor scaling.
The mixed-precision recipe#
Micikevicius et al. (2018) established the approach:
- Keep an FP32 master copy of the weights. Tiny updates ($\eta \times$ gradient) can be smaller than FP16's precision relative to the weight; accumulating them in FP16 would lose them entirely.
- Run the forward and backward passes in 16-bit for matrix multiplications and convolutions.
- Keep sensitive operations in FP32: reductions (sums, softmax, normalisation statistics), loss computation, and large accumulations.
- Loss scaling (for FP16): multiply the loss by a factor $S$ (e.g. $2^{16}$) before backpropagation so small gradients are shifted into FP16's representable range; divide gradients by $S$ before the optimiser step. Dynamic loss scaling increases $S$ periodically and halves it (skipping the step) whenever infinities or NaNs appear.
With BF16, loss scaling is generally unnecessary.
Automatic mixed precision in PyTorch#
import torch
device = "cuda"
model = MyModel().to(device)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
use_bf16 = torch.cuda.is_bf16_supported()
dtype = torch.bfloat16 if use_bf16 else torch.float16
scaler = torch.amp.GradScaler("cuda", enabled=not use_bf16) # loss scaling only for fp16
for xb, yb in loader:
xb, yb = xb.to(device, non_blocking=True), yb.to(device, non_blocking=True)
with torch.autocast(device_type="cuda", dtype=dtype):
logits = model(xb)
loss = torch.nn.functional.cross_entropy(logits, yb)
opt.zero_grad(set_to_none=True)
scaler.scale(loss).backward()
scaler.unscale_(opt) # so clipping sees true gradients
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(opt) # skips the step if inf/NaN found
scaler.update()autocast chooses the precision per operation automatically: matmuls and convolutions in 16-bit, softmax, layer norm and losses in FP32. (With enabled=False, the scaler's methods become pass-throughs, so the same loop works for BF16.)
Where the time goes: other efficiency techniques#
Mixed precision is one lever among several. Profile first (torch.profiler, NVIDIA Nsight) to find the real bottleneck.
- Keep the GPU fed: slow data loading is the most common bottleneck. Use multiple
num_workers,pin_memory=True, prefetching, and preprocess data offline where possible. - Larger batches use hardware more efficiently (with appropriate learning-rate scaling). Gradient accumulation simulates large batches when memory is limited:
accum = 4
for i, (xb, yb) in enumerate(loader):
with torch.autocast("cuda", dtype=dtype):
loss = criterion(model(xb), yb) / accum
scaler.scale(loss).backward()
if (i + 1) % accum == 0:
scaler.step(opt); scaler.update(); opt.zero_grad(set_to_none=True)- Tensor-friendly shapes: dimensions that are multiples of 8 (or 64) map better onto Tensor Cores.
- Compilation:
torch.compile(model)fuses operations and reduces Python overhead. - Fused and memory-efficient kernels: FlashAttention computes attention without materialising the full attention matrix; fused optimisers and fused layer norms reduce memory traffic.
- Activation checkpointing: recompute activations in the backward pass instead of storing them โ trading compute for memory to fit larger batches or models.
- Channels-last memory format for CNNs on modern GPUs.
- Avoid hostโdevice synchronisation inside the loop: calling
.item()or printing tensors every step forces the CPU to wait for the GPU.
Memory budgeting#
For training with Adam in mixed precision, a rough per-parameter memory cost is: 2 bytes (16-bit weights) + 4 bytes (FP32 master weights) + 8 bytes (two FP32 Adam moments) + 2โ4 bytes (gradients) โ 16โ18 bytes per parameter, plus activations. A 1-billion-parameter model therefore needs about 16โ18 GB just for weights, gradients and optimiser state โ before activations. This arithmetic explains why large models require memory-sharding techniques (next lecture).