Matmul precision#
Accelerator hardware offers several ways to compute matrix multiplication,
trading accuracy for speed: true float32 arithmetic, TensorFloat32 on NVIDIA
tensor cores, one or several bfloat16 passes on TPU, various float8 modes,
and so on.
By default, JAX leans toward speed: float32 dot products may be computed with
reduced-precision arithmetic internally (bfloat16 on TPU, TF32 on recent
GPUs). But that’s just a default, and you can take explicit control, per
operation and globally.
The precision argument, accepted by jax.lax.dot_general(),
jax.lax.dot(), and the jax.numpy functions built on them
(jnp.dot, jnp.matmul, @, jnp.einsum, convolutions), is how you
control this. It accepts two kinds of value: a dot algorithm, which names
the computation exactly, or a coarse-grained Precision
level, the older interface.
Dot algorithms#
The most direct way to control a dot product is to name the algorithm that
computes it, by passing a jax.lax.DotAlgorithmPreset (or its name as
a string) as precision:
import jax.numpy as jnp
from jax import lax
y = jnp.dot(x, x, precision="F32_F32_F32") # true float32
y = jnp.dot(x, x, precision="BF16_BF16_F32") # bf16 inputs, f32 accumulation
y = jnp.dot(x, x, precision=lax.DotAlgorithmPreset.TF32_TF32_F32) # TF32 tensor cores
Preset names follow the pattern LHS_RHS_ACCUM: the element types the
left- and right-hand sides are rounded to, and the type used for
accumulation. The available presets:
DEFAULT— an algorithm is selected based on input and output types.F32_F32_F32,F64_F64_F64— plain full-precision arithmetic.F16_F16_F16,F16_F16_F32— half-precision inputs, accumulating in half or single precision.BF16_BF16_BF16,BF16_BF16_F32— likewise forbfloat16.BF16_BF16_F32_X3,_X6,_X9— the_Xsuffix means the algorithm uses that manybfloat16operations to emulate higher precision:_X3approachesfloat32accuracy,_X6and_X9exceed it, at proportionally higher cost.TF32_TF32_F32,TF32_TF32_F32_X3— TensorFloat32, and its 3-operation higher-precision emulation.ANY_F8_ANY_F8_F32,ANY_F8_ANY_F8_F32_FAST_ACCUM— anyfloat8input types, accumulating intofloat32; theFAST_ACCUMvariant uses faster, less accurate accumulation (e.g. cuBLASLt’s fast-accumulation mode).ANY_F8_ANY_F8_ANY,ANY_F8_ANY_F8_ANY_FAST_ACCUM— as above, with the accumulation type controlled bypreferred_element_type.
Some properties of this interface:
Any input dtypes are accepted. JAX inserts casts so that the operands reach the hardware in the algorithm’s storage types: you can pass
float32arrays withprecision="BF16_BF16_F32"and the rounding is handled for you.The output type matches the inputs (under the usual promotion rules), regardless of the algorithm’s internal accumulation type, so switching algorithms doesn’t ripple type changes through your program. To instead keep the accumulator’s type, use
preferred_element_type:x16 = jnp.ones((4, 4), jnp.float16) jnp.dot(x16, x16, precision="F16_F16_F32") # f16 result jnp.dot(x16, x16, precision="F16_F16_F32", preferred_element_type=jnp.float32) # keep the f32 accumulator
Autodiff carries the same algorithm through to the backward pass. The transposed dots in the gradient computation carry the same
precisionargument as the primal. You can see this in the jaxpr of a gradient:def loss(x, w): return jnp.sum(jnp.dot(x, w, precision='BF16_BF16_F32')) x, w = jnp.ones((4, 8)), jnp.ones((8, 2)) print(jax.jit(jax.grad(loss, argnums=1)).trace(x, w).jaxpr)
{ lambda ; a:f32[4,8] b:f32[8,2]. let c:f32[4,2] = dot_general[ dimension_numbers=(([1], [0]), ([], [])) precision=BF16_BF16_F32 preferred_element_type=float32 ] a b _:f32[] = reduce_sum[axes=(0, 1) out_sharding=None] c d:f32[4,2] = broadcast_in_dim 1.0:f32[] e:f32[2,8] = dot_general[ dimension_numbers=(([0], [0]), ([], [])) precision=BF16_BF16_F32 preferred_element_type=float32 ] d a f:f32[8,2] = transpose[permutation=(1, 0)] e in (f,) }
Both
dot_generals, the forward one and the transposed one that computes the gradient, carryprecision=BF16_BF16_F32. (If you need a different backward-pass algorithm, express that withjax.custom_vjp().)Support is platform-dependent, and checked at compile time. Requesting an algorithm the backend can’t provide is a compile-time error. For example,
precision="F16_F16_F32"on CPU fails withThe precision 'F16_F16_F32' is not supported by dot_general on CPU.
If no preset fits, you can specify a fully custom algorithm with
jax.lax.DotAlgorithm, choosing the operand precision types,
accumulation type, and the number of decomposed operations directly.
Precision: the classic three levels#
The precision argument also accepts the older, coarser
jax.lax.Precision levels, which say how precise rather than
which algorithm, with device-dependent meanings. They affect only
float32 computations, and have no effect on CPU:
Precision.DEFAULT(aliases'default','fastest', andNone): fastest, least accurate. On TPU, computes inbfloat16; on GPU, uses TF32 where available.Precision.HIGH(aliases'high','bfloat16_3x','tensorfloat32'): slower, more accurate. On TPU, threebfloat16passes; on GPU, TF32.Precision.HIGHEST(aliases'highest','float32'): slowest, most accurate. On TPU, sixbfloat16passes; on GPU, truefloat32.
jnp.dot(x, x, precision='highest') # give me real float32, whatever it costs
These remain widely used, but when you care about the exact numerics (for reproducibility, for cross-platform agreement, or for f8/f16 throughput), prefer naming a dot algorithm, which pins down the computation rather than a device-dependent accuracy level.
Setting a default globally#
To change the default for every dot-like operation that doesn’t specify its
own precision, use the jax_default_matmul_precision config, as a context
manager, a config update, or an environment variable. It accepts the same
values as the precision argument, including dot algorithm preset names:
# scoped:
with jax.default_matmul_precision('highest'):
result = f(x)
# process-wide:
jax.config.update('jax_default_matmul_precision', 'BF16_BF16_F32_X3')
JAX_DEFAULT_MATMUL_PRECISION=highest python train.py
A common recipe when debugging suspected numerics problems: run once under
jax.default_matmul_precision('highest'). If the discrepancy disappears, it
was reduced-precision matmul accumulation rather than a bug.
Precision is not dtype#
Finally, a distinction: everything on this page controls how dot products
are computed for given inputs. That’s separate from the choice of dtype
your data is stored in (Arrays and jax.numpy covers defaults and
Type promotion semantics the promotion rules). Storing model parameters
or activations in bfloat16 changes memory footprint and bandwidth
everywhere; precision changes arithmetic inside individual operations.
Performance work on accelerators usually involves deciding both.