Model and Inference Engineering · Staff
Every FSDP rank clipped its gradients. Why was the global update still too large?
The question
Interview question
A team moves training from replicated data parallelism to fully sharded data parallelism. The code still calls a gradient-norm clipping utility on parameters visible to each rank. Logs show each rank below the configured norm limit. Yet some steps are much larger than the old run. Can all those local logs be true and the update still violate the intended global limit?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
Yes. For a disjoint parameter partition, the global L2 norm is the square root of the sum of squared local shard norms. Suppose four ranks each own a distinct gradient shard with norm 0.8 and the intended full-model limit is 1. Each rank sees 0.8 and does nothing. The full norm is sqrt(4 × 0.8²) = 1.6, so the intended global rule would scale the gradients down. This example assumes each parameter's gradient is counted once. Replicated parameters, nested sharding and data-parallel groups require their own accounting. PyTorch's FSDP documentation distinguishes sharded gradient clipping from the ordinary utility, while its FSDP2 tutorial notes that DTensor-based clipping can handle the sharded representation. Which API is right depends on the actual stack and version.
I would compare one step against a small unsharded reference after gradient accumulation and any mixed-precision unscale, before the optimizer update. Log the pre-clip global norm, clip coefficient, post-clip norm and whether any rank had a nonfinite value. Make sure all ranks agree on the single coefficient for the same logical update. Clipping before accumulation changes the objective, as The gradient norm is clipped on every microbatch. Why is the accumulated norm still too large? discusses. Clipping scaled gradients before unscale also changes the effective threshold. For a sharded implementation, the reduction must cover the right process group and count shared or replicated gradients according to the reference model.
The pushback is that optimizer adaptive moments might make parameter-update norm different from gradient norm anyway. Correct. Gradient clipping limits the gradient vector under a chosen norm. It does not directly guarantee a bound on the Adam parameter update. We still need to implement the claimed gradient rule correctly, then monitor update norms separately if they matter. Comparing two training runs only by final loss can miss occasional extreme steps, especially on rare data slices.
Test a fixture where every shard is individually below the threshold but the combined norm is above it, plus one with a nonfinite shard. Compare clipped gradients and the next weights against the unsharded reference. Ranks process different token counts. Why does averaging their losses change the objective? handles incorrect loss weighting when ranks have different token counts. This question asks whether the norm itself was computed across the full sharded parameter vector.
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 →