Floating-point formats: fp32, fp16, bf16
Be able to explain the difference between the precisions and when mixed precision is safe.
Prerequisites
- BBinary numbers and bitsrequired
- ENumerical stability and floating pointrequired
Intuition
A floating-point number consists of three parts: the sign, the exponent (how large) and the mantissa (how exact).
| Format | Bits | Sign | Exponent | Mantissa | Largest number | Decimal digits |
|---|---|---|---|---|---|---|
| fp32 | 32 | 1 | 8 | 23 | ~3.4·10³⁸ | ~7 |
| fp16 | 16 | 1 | 5 | 10 | 65 504 | ~3 |
| bf16 | 16 | 1 | 8 | 7 | ~3.4·10³⁸ | ~2 |
| fp8 (E4M3) | 8 | 1 | 4 | 3 | 448 | ~1 |
The decisive row is bf16. It has the same exponent as fp32 and sacrifices mantissa instead. The consequence is that it covers the same range of numbers — no overflows — but at a coarser resolution.
And that turned out to be exactly the right trade-off for deep learning: range matters more than precision. Gradients can be very small or very large, but they rarely need seven decimals of accuracy.
Formal
Why fp16 is awkward. With 5 exponent bits the smallest normal number is about and the largest 65 504. Gradients in deep networks routinely fall below that lower bound and become zero.
That is why loss scaling is required with fp16: multiply the loss by a large factor (2¹⁶, say) before the backward pass, so that the gradients land within the representable range, and divide back before the optimiser step.
With bf16 that is not needed — the exponent range is the same as fp32's. That is the whole reason bf16 has taken over where the hardware supports it.
Operations that have to be done in higher precision:
| Operation | Why |
|---|---|
| Accumulating sums | thousands of additions of small numbers — cancellation |
| Softmax | exponentiation and normalisation |
| Layer/BatchNorm | the variance of small differences |
| The optimiser's state | small updates accumulate over thousands of steps |
| Computing the loss | the same reason |
That is why «mixed precision» is mixed: the matrix multiplications are done in 16 bits, everything sensitive in 32.
The sneakiest effect is «swamping»: add a small number to a large one and it disappears entirely. In bf16 with 7 mantissa bits, 1.0 + 0.003 == 1.0. An optimiser that accumulates small updates in low precision therefore loses them silently. Near 1 the spacing is 0.0078125 and half the spacing is 0.00390625; 1.0 + 0.004 instead rounds to 1.0078125. This is why the master weights are commonly kept in fp32 even in bf16 training.
Hardware support decides what is practically possible:
| Format | Support |
|---|---|
| fp16 | NVIDIA Volta and later |
| bf16 | NVIDIA Ampere (A100) and later, TPUs, modern CPUs |
| fp8 | NVIDIA Hopper (H100) and later |
A practical rule: use bf16 if the hardware supports it, otherwise fp16 with loss scaling. fp8 for inference and, increasingly, for training very large models.
Code
import numpy as np, torch
for name, dt in (("fp32", torch.float32), ("fp16", torch.float16), ("bf16", torch.bfloat16)):
i = torch.finfo(dt)
print(f"{name}: max {i.max:.3e} min normal {i.tiny:.3e} eps {i.eps:.3e}")
# fp32: max 3.403e+38 min normal 1.175e-38 eps 1.192e-07
# fp16: max 6.550e+04 min normal 6.104e-05 eps 9.766e-04
# bf16: max 3.390e+38 min normal 1.175e-38 eps 7.812e-03
# ↑ the same range as fp32, but coarser steps
# Overflow in fp16, not in bf16
x = torch.tensor(70000.0)
print(x.half(), x.bfloat16()) # tensor(inf, dtype=float16) tensor(70144., ...)
# Underflow: a typical gradient in fp16
g = torch.tensor(1e-8)
print(g.half(), g.bfloat16()) # tensor(0., dtype=float16) tensor(1.0012e-08, ...)
# ↑ the gradient disappears entirely in fp16 — hence the need for loss scaling
# Loss scaling: scale up, compute, scale down
SCALE = 2 ** 16
print((g * SCALE).half(), (g * SCALE).half().float() / SCALE)
# tensor(0.0007, dtype=float16) tensor(9.9972e-09) ← rescued
# Swamping: small numbers disappear in an addition
for name, dt in (("fp32", torch.float32), ("bf16", torch.bfloat16)):
a = torch.tensor(1.0, dtype=dt)
b = torch.tensor(0.003, dtype=dt)
print(f"{name}: 1.0 + 0.003 = {(a + b).item()}")
# fp32: 1.0 + 0.003 = 1.003000020980835
# bf16: 1.0 + 0.003 = 1.0 ← the addition disappeared entirely
# Accumulation: sum 1 000 small numbers
n = 1000
for name, dt in (("fp16", torch.float16), ("bf16", torch.bfloat16), ("fp32", torch.float32)):
s = torch.zeros((), dtype=dt)
for _ in range(n):
s += torch.tensor(0.001, dtype=dt)
print(f"{name}: 1000 × 0.001 = {s.item():.4f} (exactly 1.0000)")
# ↑ this is why accumulation is always done in fp32
Mastery means
- Explains the exponent and the mantissa
- Compares fp32, fp16 and bf16
- Knows which operations require high precision
Sign in to do the exercises and build your mastery up.
Sources
- Goldberg — What Every Computer Scientist Should Know About Floating-Point Arithmetic — free to read
- arXiv — Mixed Precision Training — arXiv (open access; licence per article)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause