Model and Inference Engineering · Staff
FP8 training is faster until one outlier batch breaks the run. Which scale was stale?
The question
Interview question
A model trained in BF16 is moved to an FP8 matrix-multiply recipe. Throughput improves and ordinary batches look fine. After a rare high-activation batch, loss jumps and does not recover. The team says FP8 just halves the bytes, so it cannot affect the optimizer. How would you find the first incorrect value?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
FP8 has a small representable range and limited precision. To use it for a tensor, a training recipe chooses a scale that maps values into that range. If the scale was estimated from earlier batches, an unusual current batch can arrive before the scale adapts. Some values may saturate or become too coarse. The next matrix product, gradient and optimizer update can then change even if optimizer state itself stays in higher precision. NVIDIA Transformer Engine's delayed-scaling documentation describes scales derived from an amax history. That is one recipe, so first confirm this run actually uses it.
I would replay the offending batch from the last clean checkpoint, with the same token IDs, masks and sequence positions. Compare BF16 and FP8 after each boundary where quantization occurs. Log the current tensor amax, the scale used, FP8 format, saturation count, underflow distribution and output error for each quantized GEMM input and weight path. Inspect forward and backward separately. If the first large difference appears before attention, do not spend days tuning the optimizer. If the FP8 tensors remain close but the gradient becomes nonfinite later, examine accumulation precision, reduction, loss scaling where applicable and the exact fused kernel. The batch itself may be corrupt, so inspect its labels and lengths too.
An amax chart for an entire model is too broad. A single projection in one layer can have a different distribution from its neighbors, and rare language or code sequences may populate the tail. Check whether scales are per tensor, per block or another granularity in the chosen recipe, and whether a scale history was restored consistently from checkpoint. Compare the scale immediately before and after the outlier. A spike in amax can also cause later ordinary values to be represented more coarsely until the history moves on. That is a different failure from the outlier itself saturating.
For a fix, I would try a more suitable scaling recipe, granularity or format, or leave the sensitive operation in a higher precision path if the throughput cost is acceptable. Change one variable at a time and rerun the same hard batches and representative training pilot. Hold a rare-cohort validation set so an overall loss average cannot approve a bad trade. If the bad step already updated weights and optimizer moments, resume from a known clean checkpoint after correcting the recipe. A later finite loss does not prove the contaminated step was harmless.
This is different from an FP8 KV cache at serving time. Here the quantized training computation changes the gradients and therefore the learned weights. The thing to locate is the first tensor whose scaled representation no longer supports the intended update, not just the first dashboard point where loss finally rises.
Continue reading
Related questions
Read beyond the question
Explore more model and inference engineering
Follow another question in this area, or search the complete Question Library.
Browse this area →Browse Question Library →