♟ aarohi gandhi
← writing

what actually breaks low precision training

i rebuilt the arithmetic of a matmul bit by bit so every decision inside it could be changed on its own. then i found out most of what i learned was already published.

Large models train on small numbers now. Weights and activations get squeezed into 16, 8, sometimes 4 bits, and when a run falls over people reach for a familiar list of suspects. Not enough exponent bits. A stale scaling factor. The wrong rounding mode.

Those explanations almost never get measured, and on a GPU they cannot be. The number format, the width of the running sum, the order the products get added in, and the way tensors are scaled all arrive as one kernel. You cannot change one and hold the rest still.

In software you can. So I rebuilt the arithmetic of a matrix multiply in Java, bit for bit, and made every one of those decisions its own setting. Then I trained small models on real text under each one.

First, catching myself

None of it means anything if the simulated formats are wrong, so I wrote them twice. Java rounds the fast way. Python redoes it in exact fractions, no floating point anywhere. A script compares the two on every value each format can hold and every pair of values, as a product and as a sum. For the 8 bit formats that is all 65,536 pairs.

It found 154 disagreements on the first run, every one of them negative zero. A positive zero times a negative number is negative zero, these formats spend a bit pattern on it, and a Python fraction has no sign on zero. Java was right and my reference was wrong. The check caught its author on day one, which is the only reason I trust anything that came after.

What the measurements said

The MX formats scale each block of 32 numbers by a shared power of two, chosen by a formula in the OCP specification. That formula puts the largest value of a block in a range the 8 bit format cannot fully reach, so the biggest values clip. In training that cost 10% of every gradient, and I could point at the cause: 27 to 30% of activations were being clipped, because the model uses tanh and tanh outputs pile up just under 1.0, the worst possible place for that rule. One extra bit of exponent headroom fixed it completely.

Then a more interesting result. I measured how much the cast changes a tensor, and separately which way it points. Ordinary 8 bit with one scale per tensor has the largest error of anything I ran, and keeps essentially all of its gradient. Block scaled 8 bit has a smaller error, six times the bias, and loses a quarter. What matters is not how big the error is but whether it points one way. Clipping the biggest element of every block is a one way error. An error thirty times larger that points nowhere in particular does nothing.

Attention was worse than the plain model on every setting I tried, and a policy that trained fine without attention diverged with it on every seed.

Then I checked it on a real GPU, and that part is worth telling properly. The clipping held up immediately, 31.7% of activations against the 27 to 30% I had measured. The divergence did not. Three seeds, all fine. I made the gradient get cast too, ran it again, still fine. At that point the honest reading was that my simulator had a bug, so I went to read my own test instead of my results.

The test was wrong. A standard linear layer in PyTorch computes its backward as an uncast gradient times the weight, and casting that result afterwards is not the same thing as casting the operand before the multiply. Real low precision training casts both operands. Mine did. The GPU version quietly did not, so it was gentler than the claim it was checking. I wrote custom layers that cast both operands of every matmul in both directions and ran it a third time.

It diverged on all three seeds, at steps 452, 494 and 553. My simulator had said 446, 453 and 497. One extra bit of exponent headroom rescued it there too.

The two failures are the useful part. Casting only the forward pass never diverges. Casting the gradient afterwards never diverges. Casting both operands of the backward matmuls diverges every time. So this is not a fact about eight bit numbers in general. It belongs to the backward matmuls, which is also where the stochastic rounding damage lived. I predicted the softmax probabilities would be the culprit, since they also sit just under 1.0. They clip a tenth as much as the tanh outputs did. I was wrong, and holding the attention matmuls in full precision did not save the run either.

The part I would rather not write

I checked the literature after the fact, which is the wrong order, and it cost me the headline.

The scaling fix is published. NVIDIA's recipe for pre training with MXFP8, from June 2025, says not to use the floor based scale from the specification because values overflow, and to round the shared scale up instead. That is exactly what I did, arrived at from the other side. My strongest result is a reproduction.

The bias result has precedent too, in work from 2022 on quantization bias mattering more than variance. Attention being the fragile part of a low precision run is an active topic with its own 2025 paper. And the one case I thought ran against the published recipe, where rounding the scale up made every four bit run diverge, turns out to be in a 2025 paper from Meta and UMass that compares floor, ceil and nearest scale rounding directly, on real models, at a scale I cannot reach.

So what is actually mine. The measurements, which are careful: clip rates, underflow rates, gradient quality against an exact gradient, and bias separated from error size. A method I call probes, which computes what a different arithmetic would have produced at a run's own weights, without touching the run. That is how I showed stochastic rounding is not biased at any single step, and that it hurts by steering training somewhere worse. The exhaustive oracle. And two claims I published and then withdrew when more seeds and better measurements did not support them.

What I take from it

Twice now the per step numbers pointed at the wrong answer. At four bits, the scale with the better gradient is the one that dies. A cast that looks unbiased and accurate at any single point can still wreck a run through where it leads training.

And check the literature first. The work is good and the measurements are real, but I would rather have known on day one which parts were already known.

The code, every result and every retraction are at github.com/aarohigandhi/drift, with the longer version of this writeup.