import torch
import time
from torch import nn
from torch.nn import functional as F
from d2l import torch as d2l5.5 Numerics: Dtypes and Mixed Precision
Every tensor has a shape, a device, and a dtype, the numeric format of its entries. Modern accelerators can multiply 16-bit matrices faster and store them in half the memory of 32-bit matrices, so training systems choose numeric formats deliberately. This section compares their range and precision, explains dtype conversion and promotion, and develops mixed-precision training, which retains 32-bit master parameters while using lower precision for suitable operations. The representation of floating-point numbers is treated in Section 26.5; here the emphasis is on library behavior and practical failure modes.
Storage dtype, compute dtype, and accumulation dtype are separate choices. The principles below are stable, but autocast eligibility, matmul defaults, loss- scaler behavior, and hardware support vary by accelerator and library release. Treat the API cells as examples for the book’s tested environment and inspect the corresponding runtime documentation when reproducing them elsewhere.
import ml_dtypes
import time
import tensorflow as tf
from d2l import tensorflow as d2limport jax
import time
from jax import numpy as jnp
from flax import nnx
import optax
from d2l import jax as d2lfrom mxnet import autograd, gluon, np, npx
import time
from mxnet.gluon import nn
from d2l import mxnet as d2l
npx.set_np()5.5.1 Floating-Point Formats
Half-precision (float16, “fp16”) uses half the storage of fp32 and is supported by fast arithmetic on many accelerators. Its largest representable number is 65504. Squaring 300, a value that can occur as an unnormalized logit or an intermediate in a variance computation, already exceeds that range:
x = torch.tensor(300.0)
x.to(torch.float16)**2, x.to(torch.bfloat16)**2(tensor(inf, dtype=torch.float16), tensor(90112., dtype=torch.bfloat16))
x = tf.constant(300.0)
tf.cast(x, tf.float16)**2, tf.cast(x, tf.bfloat16)**2(<tf.Tensor: shape=(), dtype=float16, numpy=inf>,
<tf.Tensor: shape=(), dtype=bfloat16, numpy=90112.0>)
x = jnp.array(300.0)
x.astype(jnp.float16)**2, x.astype(jnp.bfloat16)**2(Array(inf, dtype=float16), Array(90112, dtype=bfloat16))
x = np.array(300.0)
x.astype('float16')**2array(inf, dtype=float16)
fp16 overflows to inf. The second format, bfloat16 (“bf16”, brain floating point), returns 90112: wrong in the fourth digit, since the exact square is 90000, but finite. The two formats spend the same 16 bits differently. fp16 uses 5 exponent bits and 10 mantissa bits: fine-grained steps, tiny range. bf16 keeps fp32’s full 8 exponent bits and pays with a 7-bit mantissa: fp32’s range, coarse steps. The example above demonstrates both consequences.
The MXNet cell shows only the first half of the demo, because mx.np arrays come in float16, float32, and float64 and nothing else. bfloat16 exists deeper in the engine (the mxnet.amp module lists it as a cast target) but is not a storage dtype you can astype an array to, so we state the bf16 result from the format definition. The fp16 half is unchanged: 300 squared overflows to inf.
| format | sign | exponent | mantissa |
|---|---|---|---|
| fp32 | 1 | 8 | 23 |
| tf32 | 1 | 8 | 10 |
| bf16 | 1 | 8 | 7 |
| fp16 | 1 | 5 | 10 |
| fp8 (e4m3) | 1 | 4 | 3 |
Figure 5.5.1 draws the same table as bit layouts, aligned at the sign bit so the total widths compare directly: fp32 spends its 32 bits on a wide mantissa, while bf16 keeps fp32’s 8-bit exponent but pays for it with a narrow 7-bit mantissa, and fp16 inverts the trade.
finfo reports what each bit budget buys. Three numbers matter: max, the overflow threshold; tiny, the smallest normal value, below which subnormal values provide a short, progressively less precise tail before zero; and eps, the relative step size between adjacent representable values near 1.
for dtype in (torch.float32, torch.bfloat16, torch.float16):
fi = torch.finfo(dtype)
print(f'{str(dtype):15} max {fi.max:10.3e} tiny {fi.tiny:9.3e}'
f' eps {fi.eps:8.3e}')torch.float32 max 3.403e+38 tiny 1.175e-38 eps 1.192e-07
torch.bfloat16 max 3.390e+38 tiny 1.175e-38 eps 7.812e-03
torch.float16 max 6.550e+04 tiny 6.104e-05 eps 9.766e-04
for dtype in (tf.float32, tf.bfloat16, tf.float16):
fi = ml_dtypes.finfo(dtype.as_numpy_dtype)
print(f'{dtype.name:10} max {float(fi.max):10.3e}'
f' tiny {float(fi.tiny):9.3e} eps {float(fi.eps):8.3e}')float32 max 3.403e+38 tiny 1.175e-38 eps 1.192e-07
bfloat16 max 3.390e+38 tiny 1.175e-38 eps 7.812e-03
float16 max 6.550e+04 tiny 6.104e-05 eps 9.766e-04
for dtype in (jnp.float32, jnp.bfloat16, jnp.float16):
fi = jnp.finfo(dtype)
print(f'{fi.dtype.name:10} max {float(fi.max):10.3e}'
f' tiny {float(fi.tiny):9.3e} eps {float(fi.eps):8.3e}')float32 max 3.403e+38 tiny 1.175e-38 eps 1.192e-07
bfloat16 max 3.390e+38 tiny 1.175e-38 eps 7.812e-03
float16 max 6.550e+04 tiny 6.104e-05 eps 9.766e-04
for dtype in ('float32', 'float16'):
fi = np.finfo(dtype)
print(f'{dtype:10} max {fi.max:10.3e} tiny {fi.smallest_normal:9.3e}'
f' eps {fi.eps:8.3e}')float32 max 3.403e+38 tiny 1.175e-38 eps 1.192e-07
float16 max 6.550e+04 tiny 6.104e-05 eps 9.766e-04
TensorFlow has no finfo of its own (the tf.experimental.numpy one chokes on bfloat16). The numbers live one level down, in ml_dtypes, the small NumPy-extension package that defines the bfloat16 scalar type TensorFlow depends on; each tf dtype hands over its NumPy counterpart via as_numpy_dtype.
mx.np.finfo follows the array-API naming, so the smallest normal value is the field smallest_normal rather than tiny; it passes the dtype to NumPy and repackages the result. NumPy knows no bfloat16 and MXNet stores none, so the loop covers fp32 and fp16 and we state bf16’s numbers for the record: max \(3.39 \times 10^{38}\) and tiny \(1.18 \times 10^{-38}\), matching fp32’s exponent range, with eps \(2^{-7} \approx 0.0078\).
bf16 matches fp32’s tiny and has nearly the same max: its shared exponent width gives the same normal exponent range, while its shorter mantissa makes the largest finite value slightly smaller. Its eps of 0.0078 means two to three significant decimal digits. fp16 resolves about three to four digits, overflows at 65504, enters the subnormal range below about \(6.10 \times 10^{-5}\), and reaches zero only below about \(5.96 \times 10^{-8}\). The tradeoff is precision against range. In deep learning workloads, activations and gradients can span many orders of magnitude. Insufficient range produces inf, while lower mantissa precision introduces rounding error that must be evaluated for the model and task. On accelerators with native bf16 support, that exponent range often makes bf16 preferable to fp16 for training. Older devices and deployment targets may instead favor fp16 or another format.
5.5.1.1 TF32: What Happens to fp32 Matrix Multiplication
The second row of the table, tf32, is not a storage dtype; you cannot create a tensor of it. It is a compute mode of NVIDIA tensor cores (Ampere generation and later): matrix-multiply inputs are rounded to a 10-bit mantissa while keeping fp32’s exponent, and products are accumulated in fp32. Your tensors stay fp32; only the arithmetic inside the multiplication runs faster and slightly less precisely. Most of our frameworks expose a switch for it:
print(torch.get_float32_matmul_precision())
torch.set_float32_matmul_precision('high') # allow tf32 in fp32 matmuls
print(torch.get_float32_matmul_precision(),
torch.backends.cuda.matmul.allow_tf32)
torch.set_float32_matmul_precision('highest') # restore the defaulthighest
high True
print(tf.config.experimental.tensor_float_32_execution_enabled())
tf.config.experimental.enable_tensor_float_32_execution(False)
print(tf.config.experimental.tensor_float_32_execution_enabled())
tf.config.experimental.enable_tensor_float_32_execution(True) # the defaultTrue
False
print(jax.config.jax_default_matmul_precision)
with jax.default_matmul_precision('tensorfloat32'):
print(jax.config.jax_default_matmul_precision)None
tensorfloat32
The default is 'highest': fp32 matrix multiplications compute in true fp32 and the tensor-core shortcut stays off. Setting 'high' opts in (it also flips the older allow_tf32 flag; they are two views of one setting). PyTorch 1.7 through 1.11 enabled tf32 by default, and the default changed to off in 1.12. Older tutorials can therefore describe behavior that differs from an installed version; the check above reports the setting used by the current process. For many training workloads, 'high' improves tensor-core throughput without a material accuracy change because products still accumulate in fp32. Validate that choice for the model and task, and keep 'highest' for ill-conditioned linear algebra or bitwise reproduction.
The getter answers True: TensorFlow turns tf32 on by default wherever the hardware supports it, the polarity JAX shares and today’s PyTorch inverts, so an fp32 matrix multiplication on an Ampere-or-later GPU is already taking the fast path unless you say otherwise. The switch is global, not scoped: a numerically delicate block disables it, computes, and restores it, as the cell does. On the CPU this notebook runs on there are no tensor cores and the flag changes nothing beyond its own value; the distinction takes effect on the GPUs of Section 5.7. For many training workloads, the default improves throughput without a material accuracy change because products still accumulate in fp32. Validate that choice for the model and task; disable it for ill-conditioned linear algebra or bitwise reproduction.
The unset config (None) means “let the backend pick”, and on tensor-core GPUs XLA’s pick for fp32 matrix multiplications is the fast tf32 path: the opposite polarity from PyTorch, which today makes you opt in. You opt out by requesting 'float32' (or 'highest'), and the setting is a scoped context manager rather than a global flag, so a numerically delicate block can demand full precision without touching the rest of the program. Individual operations accept the same choice per call, as in jnp.dot(A, B, precision='float32'). On the CPU this notebook runs on there are no tensor cores and every setting computes plain fp32, which is why the cell above changes nothing but the config value; the distinction takes effect on the GPUs of Section 5.7. For many training workloads, tf32 improves throughput without a material accuracy change because products still accumulate in fp32. Validate that choice for the model and task; request 'float32' for ill-conditioned linear algebra or bitwise reproduction.
MXNet is the exception: no cell here, because no user-facing switch exists. Nothing in the Python API reads or sets the fp32-matmul precision; whether an fp32 matrix multiplication takes the tensor-core shortcut is decided below MXNet, by the CUDA libraries the wheel links against. If you need the guarantees the other tabs configure (true-fp32 matmuls for ill-conditioned linear algebra, or bit-for-bit reproduction), MXNet does not give you the knob; the available workaround is to do that computation in float64 or outside the framework.
5.5.1.2 Below 16 Bits
Production inference pushes further down: int8 quantization is standard for serving, and fp8 training (a 4-bit-exponent variant for the forward pass, a 5-bit-exponent variant for gradients, with per-tensor scaling) runs on H100-class hardware (Micikevicius et al. 2022). Both require calibration machinery beyond a dtype argument, so for this book 16 bits is the floor; the exercises let you inspect the fp8 format’s finfo.
TensorFlow 2.21 ships no fp8 dtype. The formats exist one level down, in the same ml_dtypes package that supplied our finfo (ml_dtypes.finfo(ml_dtypes.float8_e4m3fn) reads them), but no tf dtype wraps them; fp8 work in the TensorFlow world happens in specialized libraries.
The fp8 pair already ships in jnp (jnp.float8_e4m3fn for the forward pass, the wider-range jnp.float8_e5m2 for gradients), and jnp.finfo reads them like any other float; what still lives in specialized libraries is the per-tensor scaling that makes them trainable.
MXNet ships no fp8 dtype at any level: not in mx.np, not in the engine’s dtype tables. For this tab the 16-bit floor is not a book decision but a hard one; fp8 experiments mean leaving the framework.
5.5.2 Dtype Rules: Promotion, Parameters, and Casts
What happens when dtypes meet in one expression? For plain tensors, PyTorch promotes to the type that can represent both:
What happens when dtypes meet in one expression? For plain tensors, TensorFlow refuses to guess. There is no promotion; an operation whose inputs disagree raises on the spot:
What happens when dtypes meet in one expression? For plain arrays, JAX promotes to a common type that can represent both:
What happens when dtypes meet in one expression? For plain arrays, MXNet promotes to the type that can represent both:
x16 = torch.ones(3, dtype=torch.float16)
x32 = torch.ones(3, dtype=torch.float32)
(x16 + x32).dtype, (x16 + 1.0).dtype, (x16 + torch.arange(3)).dtype(torch.float32, torch.float16, torch.float16)
x16 = tf.ones(3, dtype=tf.float16)
x32 = tf.ones(3, dtype=tf.float32)
try:
x16 + x32
except tf.errors.InvalidArgumentError as e:
print(e.message)
(x16 + 1.0).dtypecannot compute AddV2 as input #1(zero-based) was expected to be a half tensor but is a float tensor [Op:AddV2] name:
tf.float16
x16 = jnp.ones(3, dtype=jnp.float16)
x32 = jnp.ones(3, dtype=jnp.float32)
(x16 + x32).dtype, (x16 + 1.0).dtype, (x16 + jnp.arange(3)).dtype(dtype('float32'), dtype('float16'), dtype('float16'))
x16 = np.ones(3, dtype='float16')
x32 = np.ones(3, dtype='float32')
((x16 + x32).dtype, (x16 + 1.0).dtype,
np.result_type(x16, np.ones(3, dtype='int32')))(dtype('float32'), dtype('float16'), dtype('float16'))
Mixing two float tensors upcasts to the wider one, so an fp16 pipeline with a stray fp32 tensor quietly becomes fp32 from that point on, doubling downstream memory. Python scalars and integer tensors are weak: they adopt the float tensor’s dtype instead of dragging it up, which is why sprinkling literals like x + 1.0 into low-precision code is harmless. (Mixing fp16 with bf16 promotes to fp32, since neither contains the other.)
Only the Python scalar got through: literals like x + 1.0 are cast to the tensor’s dtype, so sprinkling them into low-precision code is harmless. Everything else, float with float or float with integer, is an error at the first operation that touches both. This strictness guards against exactly the failure that promotion invites in promotion-based libraries, where an fp16 pipeline with a stray fp32 tensor quietly becomes fp32 from that point on, doubling downstream memory; in TensorFlow the stray tensor cannot travel one op. The price is that every intended conversion is spelled tf.cast.
Mixing two float tensors upcasts to the wider one, so an fp16 pipeline with a stray fp32 tensor quietly becomes fp32 from that point on, doubling downstream memory. Python scalars and integer tensors are weak: they adopt the float tensor’s dtype instead of dragging it up, which is why sprinkling literals like x + 1.0 into low-precision code is harmless. (Mixing fp16 with bf16 promotes to fp32, since neither contains the other.)
Mixing two float arrays upcasts to the wider one, so an fp16 pipeline with a stray fp32 array quietly becomes fp32 from that point on, doubling downstream memory. Python scalars are weak: they adopt the array’s dtype instead of dragging it up, which is why sprinkling literals like x + 1.0 into low-precision code is harmless. Integer arrays behave the same way, and the wheel’s source says where the rule comes from: its promotion table is annotated “PyTorch convention: the floating operand keeps its width regardless of the integer operand”. The third expression asks that table directly through np.result_type rather than running a mixed-integer kernel.
net = nn.Sequential(nn.Flatten(), nn.Linear(784, 256), nn.ReLU(),
nn.Linear(256, 10))
def param_bytes(net):
return sum(p.numel() * p.element_size() for p in net.parameters())
param_bytes(net)814120
net = tf.keras.Sequential([tf.keras.layers.Flatten(),
tf.keras.layers.Dense(256, activation='relu'),
tf.keras.layers.Dense(10)])
net(tf.zeros((1, 28, 28, 1))) # build the weights
def param_bytes(net):
return sum(w.shape.num_elements() * tf.as_dtype(w.dtype).size
for w in net.weights)
param_bytes(net)814120
net = nnx.Sequential(lambda x: x.reshape((x.shape[0], -1)),
nnx.Linear(784, 256, rngs=nnx.Rngs(0)), nnx.relu,
nnx.Linear(256, 10, rngs=nnx.Rngs(1)))
def param_bytes(model):
return sum(p.size * p.dtype.itemsize
for _, p in nnx.state(model, nnx.Param).flat_state())
param_bytes(net)814120
net = nn.Sequential()
net.add(nn.Dense(256, activation='relu'), nn.Dense(10))
net.initialize()
X = np.random.normal(size=(8, 1, 28, 28))
net(X) # materialize the weights (initialization is deferred)
def param_bytes(net):
return sum(p.data().size * p.data().dtype.itemsize
for p in net.collect_params().values())
param_bytes(net)814120
net = net.bfloat16()
X = torch.randn(8, 1, 28, 28)
try:
net(X)
except RuntimeError as e:
print(e)
print(net(X.bfloat16()).dtype, param_bytes(net))mat1 and mat2 must have the same dtype, but got Float and BFloat16
torch.bfloat16 407060
net = tf.keras.Sequential([
tf.keras.layers.Flatten(),
tf.keras.layers.Dense(256, activation='relu', dtype='bfloat16'),
tf.keras.layers.Dense(10, dtype='bfloat16')])
X = tf.random.normal((8, 28, 28, 1))
print(net(X).dtype, param_bytes(net))<dtype: 'bfloat16'> 407060
graph, params = nnx.split(net, nnx.Param)
params16 = jax.tree.map(lambda p: p.astype(jnp.bfloat16), params)
net16 = nnx.merge(graph, params16)
X = jax.random.normal(d2l.get_key(), (8, 28, 28, 1))
print(net16(X).dtype, param_bytes(net16))
print(net16(X.astype(jnp.bfloat16)).dtype)float32 407060
bfloat16
net.cast('float16')
print(net(X.astype('float16')).dtype, param_bytes(net))float16 407060
The parameter footprint halves, inference works, and the error message shows the strictness: the fp32 input was refused rather than silently converted.
The parameter footprint halves, and note what did not happen: no error, although we fed an fp32 input to a bf16 model. The strictness of the previous cell lives at the level of raw operations; a Keras layer casts floating-point inputs to its compute_dtype at the door, so by the time any op runs, the dtypes already agree. The layer’s policy decides the arithmetic, not the input, and both faces of TensorFlow serve the same end: ops refuse to guess, layers make the choice explicit at construction.
The parameter footprint halves, but look at the first output: fp32. With dtype=None the layer promoted the bf16 parameters and the fp32 input to a common type, so casting the parameters alone bought memory but not 16-bit compute; the fp32 input dragged the arithmetic straight back up. Feeding a bf16 input (second line) keeps the whole forward pass in bf16, and so does constructing the layers with dtype=jnp.bfloat16, which pins the compute format regardless of what comes in. Nothing is refused and nothing raises: in flax the dtype story is promotion plus two explicit arguments.
The parameter footprint halves and the forward pass returns fp16. The input was cast too: feed the fp32 X to the fp16 network and type inference refuses the mismatch, the same strictness as PyTorch with the error raised one level lower, by the operator. (No Flatten layer either; Gluon’s Dense flattens trailing dimensions by default.)
For inference, casting a model can halve parameter memory, but its effect on predictions must be evaluated for the model and task. For training, casting all parameters to 16 bits can suppress small updates. A single optimizer update changes a weight by roughly \(\eta \cdot g\), often a factor \(10^{-4}\) or less of the weight’s own magnitude, and adding an increment smaller than about eps times the weight rounds to no change at all. With bf16’s eps of 0.0078, small updates can round away; in fp16, small gradients can underflow first. Mixed-precision training therefore retains fp32 master weights while performing selected computations in 16 bits.
5.5.3 Mixed-Precision Training
Mixed-precision training (Micikevicius et al. 2018) splits the work: parameters stay in fp32 (the master weights, so that small updates still register), while the expensive operations of the forward and backward pass run in a 16-bit dtype.
You do not annotate anything per layer. Inside a torch.autocast context, each eligible operation consults a built-in policy: matrix multiplications and convolutions, which dominate compute and map onto tensor cores, run in the low dtype, while selected sensitive operations such as cross-entropy run in fp32. Unlisted operations run in their input dtype, so a custom reduction over a low-precision tensor still needs an explicit fp32 cast. PyTorch maintains the eligibility lists, and inputs are cast on the fly.
In Keras the split is one global switch, tf.keras.mixed_precision.set_global_policy('mixed_bfloat16'). Every layer constructed afterwards gets the policy pair fp32 variable_dtype, bf16 compute_dtype: master weights and 16-bit arithmetic, decided per layer rather than per operation. Two cautions follow from “constructed afterwards”. The policy is consulted when a layer is built, so set it before creating the model, and it is global state, so we reset it to 'float32' at the end of every cell that flips it, or it would silently reshape every model built later in this notebook. What you compute outside the layers, such as the loss, follows no policy; we cast logits to fp32 before the loss for exactly the reason the autocast frameworks pin reductions there.
In Flax NNX this split is not a context manager; it is the pair of constructor arguments from the previous section. Leave param_dtype at its fp32 default and set dtype=jnp.bfloat16, and each layer stores fp32 parameters while its matrix multiplication runs in bf16: master weights and 16-bit compute, spelled out in the layer definition. There is no built-in per-operation policy to consult, and the flip side of that explicitness is that anything you compute outside the model, such as the loss reduction, stays in whatever dtype you give it; we cast logits to fp32 before the loss for exactly the reason the policy-based frameworks pin reductions there.
Gluon predates per-operation autocasting, and its recipe is the original low-ceremony one from the first mixed-precision papers, with the two halves assigned to different objects: cast the network to fp16, and ask the optimizer to keep the fp32 master weights. The second half is one flag, multi_precision=True among the optimizer parameters. With it, the updater allocates an fp32 copy of every fp16 weight, accumulates each update into that copy in full precision, and writes the rounded fp16 version back to the network; without it, the optimizer warns that accumulating in fp16 risks poor accuracy or slow convergence. (An mxnet.amp module with per-operation lists in the PyTorch style also exists, for fp16 and bf16 targets, but this book does not exercise it.)
Figure 5.5.2 draws the resulting loop, and with it the distinction that matters: casting a whole model puts everything in one dtype, whereas mixed precision keeps fp32 master weights that a 16-bit forward and backward pass read from and write back to.
net = nn.Sequential(nn.Flatten(), nn.Linear(784, 256), nn.ReLU(),
nn.Linear(256, 10))
with torch.autocast('cpu', dtype=torch.bfloat16):
Y = net(X)
Y.dtype, net[1].weight.dtype(torch.bfloat16, torch.float32)
tf.keras.mixed_precision.set_global_policy('mixed_bfloat16')
net = tf.keras.Sequential([tf.keras.layers.Flatten(),
tf.keras.layers.Dense(256, activation='relu'),
tf.keras.layers.Dense(10)])
Y = net(X)
tf.keras.mixed_precision.set_global_policy('float32')
Y.dtype, net.layers[1].kernel.dtype(tf.bfloat16, 'float32')
net = nnx.Sequential(lambda x: x.reshape((x.shape[0], -1)),
nnx.Linear(784, 256, dtype=jnp.bfloat16,
rngs=nnx.Rngs(0)), nnx.relu,
nnx.Linear(256, 10, dtype=jnp.bfloat16,
rngs=nnx.Rngs(1)))
net(X).dtype, jax.tree.leaves(nnx.state(net, nnx.Param))[0].dtype(dtype(bfloat16), dtype('float32'))
net = nn.Sequential()
net.add(nn.Dense(256, activation='relu'), nn.Dense(10))
net.initialize()
net.cast('float16')
trainer = gluon.Trainer(net.collect_params(), 'sgd',
{'learning_rate': 0.1, 'multi_precision': True})
loss = gluon.loss.SoftmaxCrossEntropyLoss()
with autograd.record():
l = loss(net(X.astype('float16')), np.zeros(8)).mean()
l.backward()
trainer.step(1)
master = trainer._updaters[0].states[0][0] # the fp32 master copy
net[0].weight.data().dtype, master.dtype(dtype('float16'), dtype('float32'))
Compute in bf16, storage in fp32: exactly the opposite split from net.bfloat16(). The first argument names the device type the computation runs on; on a GPU you would write torch.autocast('cuda', ...) and nothing else changes.
Compute in bf16, storage in fp32: exactly the opposite split from dtype='bfloat16'. A global policy sets this split before construction, so individual layers need no policy argument. There is no device argument; the same policy means the same thing on CPU, GPU, and TPU.
Compute in bf16, storage in fp32: exactly the opposite split from casting the parameter tree, and the fp32 input no longer drags anything up because dtype pins the compute format. Note what the cell did not need: no context manager, no device argument. The same constructor arguments mean the same thing on CPU, GPU, and TPU.
Storage in fp16, master copy in fp32, with the two halves assigned to the network and optimizer. Gluon keeps fp16 parameters in the network, so the forward pass needs no machinery at all, and hides the fp32 master weights in the updater’s state; the demonstration has to take one training step first (the master copy is allocated lazily) and then reaches into that internal state to show it. The practical consequence is identical: forward and backward run in 16 bits, the update accumulates in fp32.
The following experiment compares accuracy under fp32 and 16-bit compute. It trains the same MLP on Fashion-MNIST from the same initialization and on the same batches:
data = d2l.FashionMNIST(batch_size=256)
loader = torch.utils.data.DataLoader(data.train, batch_size=256)
batches = [b for b, _ in zip(loader, range(100))]
val_batch = next(iter(torch.utils.data.DataLoader(data.val, batch_size=1024)))data = d2l.FashionMNIST(batch_size=256)
images, labels = data.train
batches = [(tf.constant(images[k:k+256, :, :, None], tf.float32) / 255,
tf.constant(labels[k:k+256], tf.int32))
for k in range(0, 100 * 256, 256)]
val_images, val_labels = data.val
val_batch = (tf.constant(val_images[:1024, :, :, None], tf.float32) / 255,
tf.constant(val_labels[:1024], tf.int32))data = d2l.FashionMNIST(batch_size=256)
images, labels = data.train
batches = [(jnp.asarray(images[k:k+256], jnp.float32) / 255,
jnp.asarray(labels[k:k+256], jnp.int32))
for k in range(0, 100 * 256, 256)]
val_images, val_labels = data.val
val_batch = (jnp.asarray(val_images[:1024, :, :, None], jnp.float32) / 255,
jnp.asarray(val_labels[:1024], jnp.int32))data = d2l.FashionMNIST(batch_size=256)
loader = gluon.data.DataLoader(data.train, batch_size=256)
batches = [b for b, _ in zip(loader, range(100))]
val_batch = next(iter(gluon.data.DataLoader(data.val, batch_size=1024)))def train(amp):
torch.manual_seed(0) # identical init and data order for both runs
net = nn.Sequential(nn.Flatten(), nn.Linear(784, 256), nn.ReLU(),
nn.Linear(256, 10))
opt = torch.optim.SGD(net.parameters(), lr=0.1)
losses = []
start = time.perf_counter()
for X, y in batches:
with torch.autocast('cpu', dtype=torch.bfloat16, enabled=amp):
loss = F.cross_entropy(net(X), y)
loss.backward()
opt.step()
opt.zero_grad()
losses.append(loss.item())
elapsed = time.perf_counter() - start
Xv, yv = val_batch
with torch.no_grad():
accuracy = (net(Xv).argmax(1) == yv).float().mean().item()
return losses, accuracy, elapseddef train(mixed):
tf.keras.utils.set_random_seed(0) # identical init for both runs
tf.keras.mixed_precision.set_global_policy(
'mixed_bfloat16' if mixed else 'float32')
net = tf.keras.Sequential([tf.keras.layers.Flatten(),
tf.keras.layers.Dense(256, activation='relu'),
tf.keras.layers.Dense(10)])
opt = tf.keras.optimizers.SGD(learning_rate=0.1)
losses = []
start = time.perf_counter()
for X, y in batches:
with tf.GradientTape() as tape:
logits = tf.cast(net(X), tf.float32)
loss = tf.reduce_mean(
tf.keras.losses.sparse_categorical_crossentropy(
y, logits, from_logits=True))
grads = tape.gradient(loss, net.trainable_variables)
opt.apply_gradients(zip(grads, net.trainable_variables))
losses.append(float(loss))
elapsed = time.perf_counter() - start
Xv, yv = val_batch
accuracy = float(tf.reduce_mean(tf.cast(
tf.argmax(net(Xv), axis=1, output_type=yv.dtype) == yv, tf.float32)))
tf.keras.mixed_precision.set_global_policy('float32')
return losses, accuracy, elapseddef train(mixed):
dtype = jnp.bfloat16 if mixed else jnp.float32
net = nnx.Sequential(lambda x: x.reshape((x.shape[0], -1)),
nnx.Linear(784, 256, dtype=dtype,
rngs=nnx.Rngs(0)), nnx.relu,
nnx.Linear(256, 10, dtype=dtype,
rngs=nnx.Rngs(1)))
# param_dtype stays fp32, so the same key gives identical init either way
optimizer = nnx.Optimizer(net, optax.sgd(0.1), wrt=nnx.Param)
@nnx.jit
def step(model, optimizer, X, y):
def loss_fn(model):
logits = model(X).astype(jnp.float32)
return optax.softmax_cross_entropy_with_integer_labels(
logits, y).mean()
loss, grads = nnx.value_and_grad(loss_fn)(model)
optimizer.update(model, grads)
return loss
losses = []
start = time.perf_counter()
for X, y in batches:
loss = step(net, optimizer, X, y)
losses.append(float(loss))
loss.block_until_ready()
elapsed = time.perf_counter() - start
Xv, yv = val_batch
accuracy = float((net(Xv).argmax(1) == yv).mean())
return losses, accuracy, elapseddef train(mixed):
npx.random.seed(0)
net = nn.Sequential()
net.add(nn.Dense(256, activation='relu'), nn.Dense(10))
net.initialize()
net(batches[0][0]) # materialize the weights, identical for both runs
if mixed:
net.cast('float16')
trainer = gluon.Trainer(net.collect_params(), 'sgd',
{'learning_rate': 0.1,
'multi_precision': mixed})
loss = gluon.loss.SoftmaxCrossEntropyLoss()
losses = []
start = time.perf_counter()
for X, y in batches:
if mixed:
X = X.astype('float16')
with autograd.record():
l = loss(net(X), y).mean()
l.backward()
trainer.step(1)
losses.append(float(l))
elapsed = time.perf_counter() - start
Xv, yv = val_batch
if mixed:
Xv = Xv.astype('float16')
accuracy = float((net(Xv).argmax(axis=1) == yv).mean())
return losses, accuracy, elapsedloss32, acc32, sec32 = train(amp=False)
loss16, acc16, sec16 = train(amp=True)
print(f'fp32: loss {loss32[-1]:.3f}, val acc {acc32:.3f}, {sec32:.2f}s')
print(f'bf16: loss {loss16[-1]:.3f}, val acc {acc16:.3f}, {sec16:.2f}s')
d2l.plot(list(range(1, 101)), [loss32, loss16], 'step', 'loss',
legend=['fp32', 'bf16 autocast'])fp32: loss 0.815, val acc 0.737, 0.37s
bf16: loss 0.815, val acc 0.738, 1.45s
loss32, acc32, sec32 = train(mixed=False)
loss16, acc16, sec16 = train(mixed=True)
print(f'fp32: loss {loss32[-1]:.3f}, val acc {acc32:.3f}, {sec32:.2f}s')
print(f'bf16: loss {loss16[-1]:.3f}, val acc {acc16:.3f}, {sec16:.2f}s')
d2l.plot(list(range(1, 101)), [loss32, loss16], 'step', 'loss',
legend=['fp32', 'mixed_bfloat16'])fp32: loss 0.712, val acc 0.761, 0.66s
bf16: loss 0.711, val acc 0.762, 0.75s
loss32, acc32, sec32 = train(mixed=False)
loss16, acc16, sec16 = train(mixed=True)
print(f'fp32: loss {loss32[-1]:.3f}, val acc {acc32:.3f}, {sec32:.2f}s')
print(f'bf16: loss {loss16[-1]:.3f}, val acc {acc16:.3f}, {sec16:.2f}s')
d2l.plot(list(range(1, 101)), [loss32, loss16], 'step', 'loss',
legend=['fp32', 'bf16 compute'])fp32: loss 0.748, val acc 0.739, 2.75s
bf16: loss 0.748, val acc 0.740, 2.52s
loss32, acc32, sec32 = train(mixed=False)
loss16, acc16, sec16 = train(mixed=True)
print(f'fp32: loss {loss32[-1]:.3f}, val acc {acc32:.3f}, {sec32:.2f}s')
print(f'fp16: loss {loss16[-1]:.3f}, val acc {acc16:.3f}, {sec16:.2f}s')
d2l.plot(list(range(1, 101)), [loss32, loss16], 'step', 'loss',
legend=['fp32', 'fp16 multi_precision'])fp32: loss 0.655, val acc 0.783, 0.35s
fp16: loss 0.655, val acc 0.782, 10.31s
The two curves lie on top of each other and their validation accuracies agree: 16-bit rounding perturbs each step slightly, so the trajectories are not bitwise identical, but they descend at the same rate to the same place. The reported time reflects this CPU-sized experiment: bf16 may be no faster, and can be slower when casting overhead dominates. The wall-clock payoff appears on accelerators with native 16-bit matrix units and in models large enough to occupy them; speedup depends on hardware, shapes, and the input pipeline (we turn to GPUs in Section 5.7). Activations still occupy half the memory. Note what mixed precision does not buy: the master weights and any Adam state remain fp32, so the parameter and optimizer terms in the memory arithmetic of Section 5.2 do not shrink. The savings are in activations and speed.
Two footnotes for this tab. The 16-bit format here is fp16, not bf16, so the agreement of the curves also leans on the next subsection: with plain SGD on a well-scaled loss the gradients stay inside fp16’s range, which is why no loss scaling was needed. And CPUs have no fast fp16 path, so the mixed run above is slower than fp32 on this machine; as for the other tabs, the speed argument is a GPU argument.
5.5.3.1 Loss Scaling for fp16
With bf16, the recipe above is complete. fp16 has one more failure mode, and it is the opposite end of the axis from the overflow that opened this section: gradient underflow. Many gradients are small. Below about \(6.10 \times 10^{-5}\) fp16 uses increasingly coarse subnormals, and below about \(5.96 \times 10^{-8}\) values round to zero:
g = torch.tensor(1e-8)
g.half(), (g * 2**16).half()(tensor(0., dtype=torch.float16), tensor(0.0007, dtype=torch.float16))
g = tf.constant(1e-8)
tf.cast(g, tf.float16), tf.cast(g * 2**16, tf.float16)(<tf.Tensor: shape=(), dtype=float16, numpy=0.0>,
<tf.Tensor: shape=(), dtype=float16, numpy=0.0006551742553710938>)
g = jnp.array(1e-8)
g.astype(jnp.float16), (g * 2**16).astype(jnp.float16)(Array(0., dtype=float16), Array(0.000655, dtype=float16))
g = np.array(1e-8)
g.astype('float16'), (g * 2**16).astype('float16')(array(0., dtype=float16), array(0.000655, dtype=float16))
The gradient vanishes, yet the same value scaled by \(2^{16}\) is perfectly representable. That is the idea of loss scaling: multiply the loss by a large constant before backpropagation, so that by linearity every gradient is scaled into representable territory, then divide the gradients by the same constant before the optimizer step.
torch.amp.GradScaler automates it, choosing the scale dynamically: start high, shrink when scaled gradients overflow to inf (skipping that step), grow back periodically.
In Keras loss scaling is the 'mixed_float16' policy plus an optimizer wrapper. Under that policy, model.compile wraps whatever optimizer you pass in tf.keras.mixed_precision.LossScaleOptimizer, which chooses the scale dynamically: start high, shrink when scaled gradients overflow to inf (skipping that step), grow back periodically. model.fit users therefore get correct fp16 training without touching anything; a custom loop applies the wrapper itself, multiplies with its scale_loss before taking gradients, and lets the wrapped optimizer unscale and skip bad steps. None of this machinery exists for 'mixed_bfloat16' because none is needed: bf16’s normal exponent range matches fp32 and avoids fp16’s narrow-range failure in ordinary training. Hence the modern default our training cell followed: prefer 'mixed_bfloat16' where the hardware supports it, and reach for 'mixed_float16' plus the scaler only on older accelerators.
optax ships no automatic scaler, and JAX code rarely misses it. The idiom grew up on TPUs where bf16 is the native 16-bit format, and bf16 needs no scaling: its normal exponent range matches fp32 and covers the magnitudes encountered in ordinary training. The recipe of this section, dtype=jnp.bfloat16 over fp32 parameters, is therefore already complete, and we deliberately skip fp16 loss scaling. If old hardware ever forces fp16 on you, the two multiplications are yours to write: scale the loss inside loss_fn, divide the gradients before optax.apply_updates.
Gluon’s Trainer has no scaler; a static loss scale is two edits you own. Multiply: (2**10 * l).backward() instead of l.backward(). Divide: trainer.step(2**10) instead of trainer.step(1), because the step’s argument rescales gradients by its inverse before the update. The scale is yours to choose, too low and small gradients still flush to zero, too high and large ones overflow; the dynamic grow-and-shrink version the other frameworks automate lives in the mxnet.amp module mentioned earlier, which is where to look if the manual recipe stops being enough.
torch.manual_seed(0)
net = nn.Sequential(nn.Flatten(), nn.Linear(784, 256), nn.ReLU(),
nn.Linear(256, 10))
opt = torch.optim.SGD(net.parameters(), lr=0.1)
scaler = torch.amp.GradScaler('cpu')
for X, y in batches[:20]:
with torch.autocast('cpu', dtype=torch.float16):
loss = F.cross_entropy(net(X), y)
scaler.scale(loss).backward() # gradients of (scale * loss)
scaler.step(opt) # unscale, check for inf/nan, then step
scaler.update()
opt.zero_grad()
print(f'loss {loss.item():.3f}, loss scale {scaler.get_scale():.0f}')loss 1.388, loss scale 65536
Keep the two failure modes straight: GradScaler exists to prevent gradient underflow; its inf/NaN check handles overflow as a side effect by skipping the bad step. bf16 normally needs no scaler, because its normal exponent range matches fp32 and avoids fp16’s narrow-range failure. Hence the modern default: on hardware with bf16 support (Ampere and later GPUs, TPUs, recent CPUs), use bf16 autocast and stop there; reach for fp16 plus GradScaler only on older accelerators.
5.5.4 Numerical Failure Modes
This section examines common sources of non-finite results and inaccurate reductions.
Trace NaNs to earlier non-finite values. NaNs often follow an overflow to inf, because arithmetic involving infinite values can produce NaN:
s = torch.tensor(60000., dtype=torch.float16) * 2 # overflows
s, s - s(tensor(inf, dtype=torch.float16), tensor(nan, dtype=torch.float16))
s = tf.constant(60000., dtype=tf.float16) * 2 # overflows
s, s - s(<tf.Tensor: shape=(), dtype=float16, numpy=inf>,
<tf.Tensor: shape=(), dtype=float16, numpy=nan>)
s = jnp.array(60000., dtype=jnp.float16) * 2 # overflows
s, s - s(Array(inf, dtype=float16), Array(nan, dtype=float16))
s = np.array(60000., dtype='float16') * 2 # overflows
s, s - s(array(inf, dtype=float16), array(nan, dtype=float16))
An overflow can occur many operations before a NaN appears in the loss, and a parameter update containing NaN makes subsequent results non-finite. Check ranges first: is any computation in fp16, and are intermediate values in the \(10^4\) regime and approaching the 65504 ceiling? Then evaluate other causes, including the learning rate.
Under autocast with GradScaler, the gradient check catches the inf in the very step it occurs and skips the update. This behavior also supports using the framework’s mixed-precision utilities during debugging.
TensorFlow can identify the first operation that produces a non-finite value. Call tf.debugging.enable_check_numerics() and execution stops with an InvalidArgumentError at the first operation whose output contains inf or NaN. The check wraps every operation, so treat it as a debugging mode, not a default.
JAX can identify the first operation that produces NaN: set jax.config.update('jax_debug_nans', True) and execution stops with an error at the first operation whose output is NaN. When the check detects a NaN, it reruns jitted code operation by operation. Use this check as a debugging mode, not as the default execution mode.
MXNet has no switch that stops at the first bad operation, so the search is manual: bisect with np.isnan(x).any() and np.isinf(x).any() probes after suspect stages, starting from the loss and walking upstream. The diagnose-in-order advice above matters all the more when the tooling will not do the walking for you.
Use a stable log-sum-exp implementation. The direct evaluation of \(\log \sum_i \exp(x_i)\) overflows long before the answer does:
x = torch.tensor([12.0, 11.0, 10.0], dtype=torch.float16)
x.exp().sum().log(), torch.logsumexp(x, dim=0)(tensor(inf, dtype=torch.float16), tensor(12.4062, dtype=torch.float16))
x = tf.constant([12.0, 11.0, 10.0], dtype=tf.float16)
tf.math.log(tf.reduce_sum(tf.exp(x))), tf.math.reduce_logsumexp(x)(<tf.Tensor: shape=(), dtype=float16, numpy=inf>,
<tf.Tensor: shape=(), dtype=float16, numpy=12.40625>)
x = jnp.array([12.0, 11.0, 10.0], dtype=jnp.float16)
jnp.log(jnp.exp(x).sum()), jax.scipy.special.logsumexp(x)(Array(inf, dtype=float16), Array(12.41, dtype=float16))
x = np.array([12.0, 11.0, 10.0], dtype='float16')
m = x.max()
np.log(np.exp(x).sum()), m + np.log(np.exp(x - m).sum())(array(inf, dtype=float16), array(12.41, dtype=float16))
The result is 12.4, but the intermediate \(e^{12}\) exceeds 65504. The subtract-the-max shift from Section 4.4 avoids this overflow. Prefer a tested library implementation when one is available.
torch.logsumexp, F.log_softmax, and F.cross_entropy all build the shift in.
tf.math.reduce_logsumexp, tf.nn.log_softmax, and Keras’s cross-entropy losses with from_logits=True all build the shift in.
jax.scipy.special.logsumexp, jax.nn.log_softmax, and optax’s cross-entropy losses all build the shift in.
mx.np ships no logsumexp, so this is the one tab where the cell writes the shift explicitly: subtract the maximum before exponentiating, then add it back outside the log. Within a model, npx.log_softmax and Gluon’s SoftmaxCrossEntropyLoss build the same shift in.
Accumulate in fp32. Long sums in a 16-bit dtype drift, because once the running total is large, each small increment falls below the spacing of representable values and is partly or wholly rounded away:
x = torch.full((10**7,), 1e-3, dtype=torch.float16)
x.cumsum(0)[-1], x.cumsum(0, dtype=torch.float32)[-1](tensor(9784., dtype=torch.float16), tensor(10004.0439))
x = tf.fill([10**7], tf.constant(1e-3, tf.float16))
(tf.cumsum(x)[-1], tf.cumsum(tf.cast(x, tf.float32))[-1],
tf.cumsum(tf.cast(x, tf.float64))[-1])(<tf.Tensor: shape=(), dtype=float16, numpy=8192.0>,
<tf.Tensor: shape=(), dtype=float32, numpy=10004.041015625>,
<tf.Tensor: shape=(), dtype=float64, numpy=10004.043579101562>)
x = jnp.full((10**7,), 1e-3, dtype=jnp.float16)
x.cumsum()[-1], x.astype(jnp.float32).cumsum()[-1](Array(1.001e+04, dtype=float16), Array(10004.044, dtype=float32))
x = np.full((10**7,), 1e-3, dtype='float16')
(x.cumsum()[-1], x.cumsum(dtype='float32')[-1],
x.cumsum(dtype='float64')[-1])(array(4., dtype=float16),
array(9780.208),
array(10004.0435791, dtype=float64))
The fp16 running total is short by about 2 percent; keeping the accumulator in fp32 while the data stays fp16 recovers the exact answer. Means over large batches, epoch-level loss totals, and variance computations all follow this pattern, and it is precisely why autocast pins reductions and normalizations to fp32. When you write your own, pass dtype=torch.float32 to the reduction.
TensorFlow’s implementation loses enough fp16 increments to land at 8192, well below the 10004.04 sum of the stored values. Its fp32 accumulation reaches 10004.04 to the displayed precision, and fp64 supplies the reference value. The result is implementation-dependent because a parallel prefix sum groups terms differently from a left-to-right loop, but the lesson is stable: the fp16 accumulator loses small increments as partial sums grow. Means over large batches, epoch-level loss totals, and variance computations all follow this pattern. Choose an accumulator wide relative to the length of the sum, and prefer tf.reduce_sum when you need only the total so the library can use an accurate reduction tree rather than materializing every prefix.
The two totals disagree in the fourth digit. The fp32 accumulation, 10004, is the exact sum of the stored values (0.001 itself rounds to the nearest fp16, which is why the answer is not 10000); the fp16 accumulation drifts above it. The drift is milder than a naive sequential loop would produce, because XLA evaluates the scan as a tree, which keeps the partial sums small, but it is drift all the same and it grows with the length of the sum. Means over large batches, epoch-level loss totals, and variance computations all follow this pattern, and it is why our training loop cast logits to fp32 before the loss. When you write your own reduction over 16-bit data, upcast first, as the second expression does.
cumsum’s dtype argument sets the accumulator while the data stays fp16. The fp16 total stalls early: once the running sum reaches a few units, fp16’s spacing exceeds 0.001 and every further increment rounds to no change, so essentially all ten million remaining terms contribute nothing. The fp32 accumulator does far better but is not immune, since near \(10^4\) fp32’s own spacing is about 0.001, the size of one increment; the fp64 accumulator returns 10004.04, the exact sum of the stored values (0.001 itself rounds to the nearest fp16, which is why the answer is not 10000). Means over large batches, epoch-level loss totals, and variance computations all follow this pattern; when you write your own reduction over 16-bit data, hand it a wide accumulator, as the second and third expressions do.
Finally, do not expect bitwise equality across numeric configurations: tf32 versus fp32, or mixed precision on versus off, differ in the last bits by design. What reproducibility you can demand, and how to get it, is the subject of Section 5.8.
5.5.5 Summary
A dtype determines range and precision. Among the 16-bit formats, fp16 provides more precision but less range, whereas bf16 retains the exponent range of fp32 with fewer significand bits.
Casting a model (net.bfloat16()) converts its parameters and is the tool for inference; training instead uses autocast, which keeps fp32 master weights and runs matrix multiplications in 16 bits under a per-operation policy. fp16 training additionally needs GradScaler to stop small gradients from underflowing to zero; bf16 does not.
TensorFlow’s raw ops never promote; a dtype mismatch raises at the first operation, and Keras layers resolve this by casting inputs to their dtype policy, the pair of storage and compute formats fixed at construction. dtype='bfloat16' sets both and casts a model for inference; set_global_policy('mixed_bfloat16') keeps fp32 master weights over bf16 compute (reset the policy when you are done, because it is global). fp16 training additionally needs the loss scaling that LossScaleOptimizer automates under 'mixed_float16'; bf16 does not.
In flax the storage and compute formats are the two constructor arguments of every layer: casting a model for inference means bf16 in both (or one tree.map over the parameters), while mixed-precision training sets dtype=jnp.bfloat16 and leaves param_dtype at fp32, combining master weights with 16-bit matrix multiplications without a context manager. On supported hardware, bf16 usually avoids the loss scaling required by fp16.
Casting a model (net.cast('float16')) converts its parameters recursively and is the tool for inference. Training keeps the cast network but adds multi_precision=True to the optimizer parameters, which maintains fp32 master weights inside the updater: 16-bit forward and backward, full-precision accumulation. The rest of the modern menu is not on offer in this archived framework: no bf16 array dtype, no tf32 switch, no fp8, and loss scaling is either two manual lines or the unexercised mxnet.amp module.
When numerical results become non-finite or inaccurate, inspect ranges, use a stable logsumexp implementation, and accumulate long sums in fp32.
5.5.6 Exercises
- Redo the memory arithmetic of Section 5.2 for mixed-precision training with Adam: fp32 master weights, fp32 gradients, two fp32 moment estimates, and bf16 activations. Which term dominates now, and how does the total compare with all-fp32 training?
- Time the fp32 and 16-bit runs of
trainagainst each other while shrinking the hidden width and the batch size. Find a model small enough that mixed precision is slower than fp32, and explain where the crossover comes from. - Print every field of
torch.finfo(torch.float8_e4m3fn)(in JAX,jnp.finfo(jnp.float8_e4m3fn); in TensorFlow,ml_dtypes.finfo(ml_dtypes.float8_e4m3fn); MXNet has no fp8 dtype, so borrow the standaloneml_dtypespackage) and compare with fp16 and bf16. Explain the namee4m3fn, including why itsmaxis 448 rather than the value the exponent bits alone would suggest. - Under autocast, normalization layers run in fp32. To see why, take the
RMSNormlayer of Section 5.4, feed it inputs with a standard deviation of about 100, and force the computation to fp16. Which intermediate quantity fails first?