The core bet: gradients in four bits
Every low-precision training scheme makes one bet: that the signal in weights, activations, and gradients survives being squeezed into a coarse grid, as long as you keep a faithful copy of the state that must not drift. FP4 makes the most aggressive version of that bet. The expensive part of training a transformer is the general matrix multiplies (GEMMs) inside the linear layers — forward, and the two in the backward pass. Those GEMMs are what FP4 targets, because that is where the FLOPs and the memory traffic live.
The catch is that 4 bits is not ‘a bit worse than 8’ — it is a different regime. With so few representable values, naive rounding destroys small updates and clips large ones, and error that would be noise in FP8 becomes bias that steers the whole run. So FP4 training keeps a high-precision master copy of the weights, does the delicate arithmetic (accumulation, the optimizer, normalization) in BF16 or FP32, and spends its budget making the 4-bit operands themselves trustworthy.
The E2M1 format: sixteen numbers
The 4-bit float used for training is E2M1: 1 sign bit, 2 exponent bits, 1 mantissa bit. The exponent bias is 2^(2-1) - 1 = 1. Working the encoding out gives the mantissa fraction as 0 or 0.5 (one bit), and four exponent fields, one of which (E=0) is subnormal:
subnormal (E=0): 2^(1-bias) * (M/2) -> {0, 0.5}
normal (E=1): 2^0 * (1 + M/2) -> {1.0, 1.5}
normal (E=2): 2^1 * (1 + M/2) -> {2.0, 3.0}
normal (E=3): 2^2 * (1 + M/2) -> {4.0, 6.0}So the positive magnitudes are {0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0}, and with the sign bit that is 16 codes total. Note what is missing: E2M1 spends every code on a finite number — there is no inf, no NaN. The largest magnitude is 6.0, the smallest nonzero is 0.5, and the gap between adjacent values is never smaller than 0.5. That grid is the entire vocabulary FP4 has for numbers.
The precision-and-range wall of four bits
Two limits fall out of that grid immediately. Range: the ratio of largest to smallest nonzero magnitude is only 6.0 / 0.5 = 12. FP8’s E4M3 spans roughly 448 / 0.0019 ≈ 240,000; FP4 spans twelve. Anything more than ~12× larger than the smallest value you care about either clips to 6.0 or flushes to zero.
Precision: with at most one mantissa bit, the step between neighbors is huge — from 4.0 to 6.0 is a single jump of 50%, so a value of 5.0 rounds to 4.0 or 6.0, a 20% error either way. Neural network weights and activations, though, are roughly bell-shaped with heavy tails: most values small and clustered, a few large outliers. A single scale for a whole tensor cannot serve both — set it for the outliers and the bulk collapses toward zero; set it for the bulk and the outliers clip. Everything that follows is a way around this one wall.
Micro-scaling: MXFP4 and NVFP4 block scales
The escape is to stop using one scale per tensor and use one scale per small block of contiguous values — micro-scaling. Each block gets a multiplier chosen for its own local magnitude, so an outlier in one block cannot wreck the resolution of a quiet block next door.
Two formats dominate. MXFP4 (the Open Compute standard) uses blocks of 32 values sharing an 8-bit E8M0 scale — a pure power of two, 2^(s-127), no mantissa. Because the scale is a power of two, applying it is an exact exponent shift with no rounding. NVFP4 (NVIDIA’s Blackwell format) goes finer: blocks of 16 with an E4M3 (8-bit float) per-block scale, plus one FP32 per-tensor scale on top. The E4M3 scale can land between powers of two, fitting each block tighter, at the cost of a second, non-exact multiply. Smaller blocks and richer scales cost more metadata but track the data more faithfully — the central FP4 trade-off.
The scaling math, worked
Quantizing a block x is: pick a scale so the block’s largest magnitude maps near FP4’s ceiling of 6.0, divide, round each element onto the E2M1 grid, and remember the scale for dequantization.
amax = max_i |x_i| # block max magnitude
scale = 2^ceil(log2(amax / 6.0)) # E8M0 power-of-two
q_i = round_to_E2M1( x_i / scale ) # onto the FP4 grid
x_hat = q_i * scale # dequantizedWorked: a block with amax = 22.0. Then 22/6 ≈ 3.67, ceil(log2 3.67) = 2, so scale = 2^2 = 4. An element x = 22 becomes 22/4 = 5.5 → rounds to 6.0, dequantized 6.0 × 4 = 24. An element x = 3.0 becomes 0.75 → 0.5, back to 2.0. The block max is preserved well; mid-range values carry visible error — which is exactly why a smaller block or a float scale, covering a narrower spread, quantizes more faithfully.
Stochastic rounding: keeping updates unbiased
Round-to-nearest has a fatal flaw for training: a weight update smaller than half the local step always rounds back to where it started, so the weight never moves and learning stalls — the coarser the grid, the worse it bites, and FP4 is very coarse. Stochastic rounding fixes it by making the rounding direction random, weighted by distance. For x between grid neighbors a < x < b:
round(x) = b with probability p = (x - a)/(b - a)
= a with probability 1 - p
E[round(x)] = a*(1-p) + b*p = a + (b-a)*p = x # unbiasedThe expected rounded value equals x exactly, so error has zero mean — it becomes noise the optimizer averages out rather than a systematic drift. A tiny update now nudges the weight to the next grid point with a small probability, and across thousands of steps those probabilistic nudges accumulate into real motion. Stochastic rounding is what lets small gradients survive a grid of sixteen points.
Hadamard rotations: spreading the outliers
Block scaling helps, but a single giant outlier still forces its whole block to a coarse scale. The trick is to rotate the data so no single coordinate is an outlier. Multiply by an orthogonal Hadamard matrix H (entries ±1/√n, with H H^T = I). A Hadamard transform mixes every coordinate into every output, so a spike concentrated in one channel gets smeared across all of them — the rotated distribution is closer to Gaussian, with a much smaller max-to-median ratio inside each block.
The reason it is free is matmul invariance. Because H H^T = I, insert H on the activations and H^T on the weights and the product is unchanged:
(X H)(H^T W) = X (H H^T) W = X WYou quantize XH and H^TW — the tamed, outlier-free versions — instead of the raw operands, and the exact GEMM result is recovered. The fast Walsh–Hadamard transform costs only O(n log n), cheap beside the matmul it protects.