Skip to content
AI-grafen
FAI engineeringDeep learning· about 90 min· fast-moving, sources checked often· verified 2026-09-21· EN

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 partThe precisionWhy
Matrix multiplicationsbf16/fp16the fastest, tolerates a coarse resolution
The activationsbf16/fp16memory
The master weightsfp32small updates must not vanish
The optimiser statefp32it accumulates over thousands of steps
Softmax, normalisation, the lossfp32numerically 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:

fp16bf16
Loss scalingrequirednot needed
The risk of overflowyesno
The precisiona better mantissacoarser
The hardwareVolta and laterAmpere and later

If the hardware supports bf16 — use bf16. It removes the whole loss scaling problem.

Formal

Loss scaling for fp16, step by step:

  1. Multiply the loss by SS (a typical starting value is 2¹⁶).
  2. Run backward() — every gradient is now SS times larger and lands within fp16's range.
  3. Unscale the gradients (divide by SS) before the clipping and the optimiser step.
  4. Detect inf/NaN in the gradients: skip the step and halve SS.
  5. If no overflows have occurred for a while: double SS.

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
1Run the same step in fp32 — if it is NaN there too the fault is not the precision
2Is the loss computation in fp32? (autocast should be excluded around it)
3Does unscale_ happen before the clipping?
4Is there a log(0), a division by zero or a sqrt of a negative number?
5Is some normalisation in fp16 with very small variances?
6Switch 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:

OutsideWhy
loss.backward()the backward pass inherits the precision from the forward pass
The optimiser stepit should run in fp32
The loss computation for sensitive lossesnumerical 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

All the sources and licences