Skip to content

Mixed Precision Training

Learning contract override: Prerequisite: a stable full-precision PyTorch loop and supported accelerator for the AMP branch. Time: 75–90 minutes for a matched comparison. Evidence: AMP status, runtime, memory, numerical checks, and held-out delta versus full precision.

What This Is

Mixed precision training uses reduced-precision and float32 operations to reduce memory use and potentially speed up training. It aims to preserve baseline quality, but numerical behavior and final metrics must be checked. Supported accelerators can run eligible reduced-precision operations substantially faster.

When You Use It

  • training large models that do not fit in GPU memory at full precision
  • speeding up training on GPUs with tensor core support (Volta, Ampere, Hopper)
  • scaling batch size within fixed memory constraints
  • training production models where wall-clock time matters

Tooling

  • torch.amp.GradScaler — scales loss to prevent gradient underflow in float16
  • torch.amp.autocast — automatically casts operations to float16 where safe
  • torch.float16 and torch.bfloat16 — the two common reduced-precision formats

How It Works

In the usual PyTorch AMP workflow, model parameters and optimizer state remain float32 while autocast selects lower or float32 precision per operation. GradScaler multiplies the loss before backpropagation so small float16 gradients are less likely to underflow before the optimizer updates the float32 parameters.

from torch.amp import GradScaler, autocast

scaler = GradScaler("cuda")

for X_batch, y_batch in train_loader:
    optimizer.zero_grad(set_to_none=True)

    with autocast(device_type="cuda"):
        logits = model(X_batch)
        loss = loss_fn(logits, y_batch)

    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

float16 vs bfloat16

Format Range Precision Best For
float16 narrow higher mantissa bits older GPUs, needs GradScaler
bfloat16 same as float32 lower mantissa bits Ampere+, often no scaler needed

If your GPU supports bfloat16, it is often simpler because the wider range means gradients rarely underflow:

with autocast(device_type="cuda", dtype=torch.bfloat16):
    logits = model(X_batch)
    loss = loss_fn(logits, y_batch)

loss.backward()
optimizer.step()

Quick Quiz

  1. What is the main benefit of mixed precision training?
    a) Higher model accuracy
    b) Reduced memory usage and faster computation
    c) Simpler code
    d) Better generalization

  2. When do you need GradScaler?
    a) Always with mixed precision
    b) Only with float16, not bfloat16
    c) Only on older GPUs
    d) Never, autocast handles it

  3. What should you keep in float32 during mixed precision?
    a) All weights
    b) Loss functions and batch norm statistics
    c) Only the optimizer
    d) Nothing, autocast handles everything

Validation Pattern

Validation also benefits from autocast for speed, but it does not need the scaler:

model.eval()
with torch.no_grad():
    with autocast(device_type="cuda"):
        logits = model(X_valid)
        val_loss = loss_fn(logits, y_valid)

What To Keep In float32

Some operations are numerically unstable in float16:

  • loss functions (autocast handles this automatically)
  • batch normalization running statistics
  • small learning rate updates
  • operations with large reductions (softmax over long sequences)

The autocast context manager handles most of these cases automatically.

Failure Pattern

Enabling float16 autocast without appropriate scaling can make small gradients underflow to zero and stall learning. NaNs or infinities more often indicate overflow or another numerical instability; GradScaler detects non-finite gradients and adjusts the scale, but does not fix every source.

Another failure: assuming mixed precision always helps. On CPUs or older GPUs without tensor cores, it may actually be slower.

Common Mistakes

  • forgetting scaler.update() after scaler.step(), which freezes the scale factor
  • using loss.backward() instead of scaler.scale(loss).backward()
  • applying gradient clipping outside the scaler workflow
  • expecting speed gains on hardware without tensor cores

Gradient Clipping With Mixed Precision

scaler.scale(loss).backward()
scaler.unscale_(optimizer)  # unscale before clipping
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()

Practice

  1. Compare training speed with and without mixed precision on the same model.
  2. Monitor memory usage with torch.cuda.max_memory_allocated() in both modes.
  3. Add gradient clipping to a mixed precision loop and verify it works correctly.
  4. Switch between float16 and bfloat16 and compare stability.
  5. Explain why the GradScaler is necessary for float16 but often unnecessary for bfloat16.

Checkpoint

  • [ ] Implement autocast and GradScaler in a training loop
  • [ ] Compare memory usage with and without mixed precision
  • [ ] Handle gradient clipping correctly with the scaler
  • [ ] Choose between float16 and bfloat16 for your hardware
  • [ ] Monitor for NaN losses and adjust scaling if needed

Runnable Example

Run the optimization and PEFT workflow from the repository root:

.venv/bin/python labs/optimization-regularization-and-peft/src/peft_workflow.py

The workflow prints amp_enabled and exercises the autocast/GradScaler branch only when supported CUDA hardware is available. Treat it as a wiring check; record runtime, memory, non-finite gradients, and held-out quality from a matched full-precision rerun before keeping AMP.

Longer Connection

Continue with PyTorch Training Loops for the full loop structure, and Optimizers and Regularization for the optimizer patterns that interact with mixed precision.

Further Reading