Model and Inference Engineering · Staff
The gradient norm is capped at one. Why does the mixed-precision run barely learn?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
Look at whether clipping happens before or after gradient unscaling. In a common fp16 AMP setup, the trainer multiplies loss by a scale S before backpropagation to avoid tiny gradients underflowing. The resulting .grad values are scaled too. If the code clips those values to norm one and then GradScaler divides by S, the actual update can have norm about 1/S, not one. A large scale turns an intended safety bound into a near-zero update. PyTorch's AMP examples explicitly call scaler.unscale_(optimizer) before gradient inspection or clipping.
Use a simple number check. Let the true accumulated gradient norm be 2, and the loss scale be 1024. Its scaled norm is 2048. Clipping that to 1 and then unscaling leaves norm about 0.00098. Correct ordering unscales to norm 2, then clips to 1. This illustration assumes ordinary global norm clipping and no other gradient transforms. It is worth measuring the real values because a changing loss scale can make the symptom intermittent.
The safe order for this setup is accumulate all intended microbatches, unscale once for the optimizer, inspect nonfinite values, clip the unscaled gradients, then let the scaler decide whether to execute the optimizer step and update its scale. Do not unscale in the middle of accumulation and keep adding scaled gradients afterward. Log the scale, norm before and after clipping, skipped-step count and parameter delta. Compare a small deterministic update to a full-precision baseline where possible.
What if the model uses bf16 and no loss scaler? Then this specific explanation may not apply. Ask which precision path and scaler are actually active. The gradient norm is clipped on every microbatch. Why is the accumulated norm still too large? covers clipping each microbatch before accumulation and Every FSDP rank clipped its gradients. Why was the global update still too large? covers clipping each shard independently. Here the completed gradient is clipped at the right time in the batch, but in the wrong units.
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 →