๐Ÿ”— Deep Learning ยท Lecture 31 of 38

Mixed-Precision Training and GPU Efficiency

Training in 16-bit arithmetic roughly halves memory and can multiply throughput. We explain float16 and bfloat16, loss scaling, automatic mixed precision, and other practical techniques to make GPUs work harder.

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#

FormatExponent bitsMantissa bitsRange (approx.)Precision (decimal digits)
FP32823$10^{-38}$ to $3 \times 10^{38}$~7
FP16510$6 \times 10^{-5}$ (normal) to 65,504~3
BF1687same as FP32~2โ€“3
FP8 (E4M3 / E5M2)4 / 53 / 2narrow~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:

  1. 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.
  2. Run the forward and backward passes in 16-bit for matrix multiplications and convolutions.
  3. Keep sensitive operations in FP32: reductions (sums, softmax, normalisation statistics), loss computation, and large accumulations.
  4. 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#

python
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.

  1. 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.
  2. Larger batches use hardware more efficiently (with appropriate learning-rate scaling). Gradient accumulation simulates large batches when memory is limited:
python
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)
  1. Tensor-friendly shapes: dimensions that are multiples of 8 (or 64) map better onto Tensor Cores.
  2. Compilation: torch.compile(model) fuses operations and reduces Python overhead.
  3. Fused and memory-efficient kernels: FlashAttention computes attention without materialising the full attention matrix; fused optimisers and fused layer norms reduce memory traffic.
  4. Activation checkpointing: recompute activations in the backward pass instead of storing them โ€” trading compute for memory to fit larger batches or models.
  5. Channels-last memory format for CNNs on modern GPUs.
  6. 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).

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ”— Deep Learning

PyTorch Fundamentals: Tensors, Autograd, Modules and the Training Loop

A practical tour of PyTorch โ€” tensors and devices, autograd, nn.Module, Dataset and DataLoader, optimisers, and a complete, correct training and evaluation loop you can reuse in every project.

Beginnerโฑ 5 min#124
๐Ÿ”— Deep Learning

Debugging Neural Network Training: A Systematic Recipe

Neural networks fail silently โ€” they train, but badly. We present a systematic recipe for finding bugs, from data inspection and overfitting a single batch to monitoring activations, gradients and learning curves.

Intermediateโฑ 6 min#126
๐Ÿ”— Deep Learning

Distributed Training: Data, Model, Pipeline and Sharded Parallelism

Large models and datasets need many accelerators. We explain data parallelism with all-reduce, sharded data parallelism (ZeRO/FSDP), tensor and pipeline model parallelism, and how they combine at scale.

Advancedโฑ 6 min#128