Mixed precision training
Be able to train in bf16/fp16 with loss scaling and know when it goes wrong.
Prerequisites
Intuition
Mixed precision means that different parts of the training run at different precision:
| The part | The precision | Why |
|---|---|---|
| Matrix multiplications | bf16/fp16 | the fastest, tolerates a coarse resolution |
| The activations | bf16/fp16 | memory |
| The master weights | fp32 | small updates must not vanish |
| The optimiser state | fp32 | it accumulates over thousands of steps |
| Softmax, normalisation, the loss | fp32 | numerically sensitive |
The gain is roughly 2–3× faster training and nearly halved activation memory, with in practice unchanged quality.
Two ways, and they differ in how much trouble they cause:
| fp16 | bf16 | |
|---|---|---|
| Loss scaling | required | not needed |
| The risk of overflow | yes | no |
| The precision | a better mantissa | coarser |
| The hardware | Volta and later | Ampere and later |
If the hardware supports bf16 — use bf16. It removes the whole loss scaling problem.
Formal
Loss scaling for fp16, step by step:
- Multiply the loss by (a typical starting value is 2¹⁶).
- Run
backward()— every gradient is now times larger and lands within fp16's range. - Unscale the gradients (divide by ) before the clipping and the optimiser step.
- Detect
inf/NaNin the gradients: skip the step and halve . - If no overflows have occurred for a while: double .
That is called dynamic loss scaling and is handled by torch.amp.GradScaler. Point 3 is the one most often got wrong — if you clip the gradients before unscale_ you clip against a norm 65 536 times too large, and the clipping does nothing.
The debugging order when mixed precision gives NaN:
| # | The check |
|---|---|
| 1 | Run the same step in fp32 — if it is NaN there too the fault is not the precision |
| 2 | Is the loss computation in fp32? (autocast should be excluded around it) |
| 3 | Does unscale_ happen before the clipping? |
| 4 | Is there a log(0), a division by zero or a sqrt of a negative number? |
| 5 | Is some normalisation in fp16 with very small variances? |
| 6 | Switch to bf16 — it solves the overflow problems entirely |
Check 1 first, always. Mixed precision often gets the blame for faults that are there in fp32 too.
What should not be inside autocast:
| Outside | Why |
|---|---|
loss.backward() | the backward pass inherits the precision from the forward pass |
| The optimiser step | it should run in fp32 |
| The loss computation for sensitive losses | numerical stability |
Inference in low precision is simpler: no optimiser, no gradients. bf16 or fp16 straight off nearly always works, and from there the road goes on to int8 and int4 with quantisation.
Always measure. Run 200 steps in fp32 and 200 in mixed precision and compare the loss curves. If they lie on top of each other everything is fine; if they diverge there is a problem that does not show up in a single measurement point.
Code
import torch, torch.nn as nn
# bf16 — the simplest, no scaler needed
def train_bf16(model, dataloader, opt, epochs=1):
for _ in range(epochs):
for x, y in dataloader:
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad(set_to_none=True)
loss.backward() # OUTSIDE autocast
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step()
# fp16 — requires a GradScaler, and the order is critical
def train_fp16(model, dataloader, opt, epochs=1):
scaler = torch.amp.GradScaler("cuda")
for _ in range(epochs):
for x, y in dataloader:
with torch.autocast("cuda", dtype=torch.float16):
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad(set_to_none=True)
scaler.scale(loss).backward() # 1-2: scale and go backwards
scaler.unscale_(opt) # 3: MUST come before the clipping
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(opt) # 4: skips the step at inf/NaN
scaler.update() # 5: adjusts the scale factor
# Keep a sensitive block in fp32
with torch.autocast("cuda", dtype=torch.bfloat16):
h = model.backbone(x)
with torch.autocast("cuda", enabled=False): # off for just this
logits = model.head(h.float())
loss = my_sensitive_loss(logits, y)
# Compare the curves instead of trusting that it works
def compare_precision(build_model, dataloader, steps=200):
curves = {}
for name, dtype in (("fp32", None), ("bf16", torch.bfloat16), ("fp16", torch.float16)):
torch.manual_seed(0)
m = build_model()
opt = torch.optim.AdamW(m.parameters(), lr=1e-4)
scaler = torch.amp.GradScaler("cuda", enabled=(dtype is torch.float16))
losses = []
for i, (x, y) in enumerate(dataloader):
if i >= steps:
break
with torch.autocast("cuda", dtype=dtype, enabled=dtype is not None):
loss = nn.functional.cross_entropy(m(x), y)
opt.zero_grad(set_to_none=True)
scaler.scale(loss).backward()
scaler.step(opt); scaler.update()
losses.append(float(loss))
curves[name] = losses
return curves
# Detect NaN early instead of a hundred steps later
def nan_guard(model):
def hook(name):
def f(m, i, o):
t = o[0] if isinstance(o, tuple) else o
if torch.is_tensor(t) and not torch.isfinite(t).all():
raise RuntimeError(f"non-finite values in {name}")
return f
for name, m in model.named_modules():
m.register_forward_hook(hook(name))
nan_guard is worth having during development. Without it a NaN is noticed only when the loss becomes NaN, and then it is hard to know which layer caused it.
Mastery means
- Sets mixed precision training up
- Uses loss scaling correctly
- Debugs NaN and divergence
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Mixed Precision Training — arXiv (open access; licence per article)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- NVIDIA — Mixed Precision Training Guide — documentation, free to read