Model and Inference Engineering · Principal
One training rank has nonfinite gradients while the global loss looks normal. Can the step commit?
The question
Interview question
At the end of a gradient accumulation window, rank 17 reports a nonfinite gradient in one layer. The dashboard's global training loss looks normal. Some ranks have entered their optimizer step and some have not. The run is expensive and the team suggests skipping only rank 17. What do you do?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
I would halt the step boundary. Replicas cannot quietly make different weight updates and then claim to be one training run. If rank 17 skips while other data parallel replicas step, their parameters and optimizer moments can diverge. With sharded or model parallel state, the situation can be worse because different participants may hold different pieces of what was meant to be one logical update. The finite scalar on the dashboard is not a certificate that every gradient or optimizer state is finite. It may have been logged before backward, averaged over other samples, or aggregated along a different path.
First I want the precise point of the failure. Was the loss finite on rank 17 before backward? Which parameter or shard became nonfinite, before or after gradient synchronization, unscaling and clipping? What precision is used for forward, gradients, master weights and optimizer math? Was this a single microbatch in an accumulation window or the completed effective batch? In ordinary data parallel training, gradients are averaged across ranks before a step, so a nonfinite value that reaches the reduction will usually affect peers too. An apparently rank local report may be taken before that reduction, involve an unsynchronized accumulation window, a sharded component, or simply reflect missing telemetry. I would not assume one bad rank is safely isolated.
For an FP16 run with gradient scaling, overflow can occur even when the forward loss is finite. PyTorch's AMP examples explain that GradScaler.step checks unscaled gradients and skips an optimizer step when it finds inf or NaN. They also put the check and scale update at the effective batch boundary during accumulation. That library behavior is local to the optimizer being stepped. Our distributed training protocol still has to decide whether all participants will skip this logical step. BF16 training may not use a gradient scaler at all, so the diagnosis cannot start and end with lowering a scale.
The safe path is a coordinated decision before any participant changes model or optimizer state. Gather a finite or nonfinite signal from every relevant rank and optimizer, decide for the full logical update, and either have everyone step or have everyone skip. Zero gradients only after that decision. Keep the scheduler, EMA if used, optimizer step counters, loss scale and consumed data accounting consistent with the chosen policy. A skipped update is not step 80,001 merely because the data loader advanced. Record whether those samples will be replayed or deliberately counted as consumed. That choice affects reproducibility and learning rate schedules driven by steps or tokens.
The scenario says some ranks already entered the optimizer step. Then I cannot assume a late all reduce repairs it. Inspect whether they actually mutated weights or moments. If any did, stop the group and restore from the last coherent checkpoint at a completed update boundary. Do not average the replicas back together and hide a partially committed update. If the engine has a genuinely staged update that can be rolled back for every participant, use its tested protocol. Most ordinary training loops do not have that guarantee. The checkpoint's committed update boundary matters here, because this step may have split before the next save exists.
After containment, reproduce the offending microbatch and inspect its token IDs, labels, attention mask and sequence lengths. Compare the failing kernel or fused operation with a higher precision reference on a small fixture. Check for an excessive loss scale, a bad data row, unstable softmax or normalization, and an optimizer state that was already corrupted on restore. Gradient clipping after nonfinite values is not a cure. If the interviewer says this happens once every few hundred steps only in FP16 and all ranks skip together, a scaler adjustment and monitored continuation may be reasonable. If the same sample repeatedly breaks BF16 or FP32, I would block the run until the data or numerical bug is understood. The invariant is the same in both cases: one logical update has one outcome across its participants.
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 →