Gradient clipping is a single line in most training loops and one of the least understood. Everyone copies clip_grad_norm_(model.parameters(), 1.0) from somewhere. Few people can say what the 1.0 refers to, where the call has to sit relative to mixed precision and gradient accumulation, or why the same number behaves differently once the model is sharded across 64 GPUs. Clipping exists to stop a single bad batch from throwing the parameters somewhere the model never recovers from. It only does that job if it measures the right quantity at the right moment.
This page builds clipping from the arithmetic up. It covers the difference between clipping by norm and clipping by value, and the exact order of operations in a step that uses AMP and accumulation. It explains how the global norm must be assembled under data, tensor and pipeline parallelism, what clipping does to Adam's moment estimates, and how to pick and audit a threshold. Finally it treats the clipping statistics as what they really are: one of the cheapest training-health signals you have.
What clipping computes
Collect every parameter gradient into one long vector g. Its L2 norm is the square root of the sum of squares of every element, across every tensor. Norm clipping computes coef = max_norm / (norm + eps) and, if coef < 1, multiplies every gradient by it. The direction of the update is preserved exactly; only its length is capped. When the norm is under the threshold nothing happens at all.
Value clipping (clip_grad_value_) is a different operation. It clamps each element independently into [-v, v]. That changes the direction of the gradient, because large components are cut while small ones survive untouched. It has niche uses, for example taming a few exploding embedding rows, but for large-model training norm clipping is the default and the rest of this page means norm clipping.
Two details matter. First, the norm is computed over all parameters together, not per tensor. Per-tensor clipping exists (and adaptive variants clip per unit, see below) but it is a different algorithm with a different threshold scale. Second, the threshold is an absolute number in gradient units. It depends on model size, loss normalization (mean or sum over tokens), and anything that rescales the loss. A threshold tuned with a per-token mean loss is meaningless if someone switches to a summed loss.
Where it sits in the step
Three operations commonly end up in the wrong order. Under AMP with fp16, the loss is multiplied by a large scale factor before backward, so the raw gradients are inflated by that factor. Clipping them before scaler.unscale_(optimizer) compares a scaled norm against an unscaled threshold, which clips almost every step. With gradient accumulation, clipping inside the micro-batch loop caps each partial sum instead of the total. And under DDP, gradients are averaged during backward, so by the time backward returns on the last micro-batch they are already global; that is why plain DDP needs no extra communication for the norm.
scaler = torch.amp.GradScaler("cuda") # only needed for fp16; bf16 skips it
for step, batch_group in enumerate(loader):
optimizer.zero_grad(set_to_none=True)
for i, mb in enumerate(batch_group): # K micro-batches
ctx = model.no_sync() if i < len(batch_group) - 1 else nullcontext()
with ctx, torch.autocast("cuda", dtype=torch.float16):
loss = model(**mb).loss / len(batch_group)
scaler.scale(loss).backward()
scaler.unscale_(optimizer) # 1. real gradient units
total = torch.nn.utils.clip_grad_norm_( # 2. global norm, then scale
model.parameters(), max_norm=1.0)
scaler.step(optimizer) # 3. skipped if inf/nan was found
scaler.update()
clipped = float(total) > 1.0
log(step=step, grad_norm=float(total), clipped=clipped)The return value of clip_grad_norm_ is the norm measured before clipping. That is the number to log. If you log the norm after clipping you will see a flat line at the threshold during exactly the events you most need to see. With bf16 autocast there is no loss scaler; drop the scaler lines and keep the same order.
Non-finite gradients deserve their own rule. If any gradient is inf or NaN, the norm is too, the coefficient becomes zero or NaN, and multiplying turns the whole gradient into NaN. With the fp16 scaler this is handled for you: unscale_ records the overflow and scaler.step skips the update. Without a scaler, pass error_if_nonfinite=True or check torch.isfinite(total) yourself and skip the step, because one NaN update destroys the run.
The global norm under parallelism
Once parameters stop being fully replicated, "the norm" needs care. The quantity you want is the norm of the full logical gradient, as if one GPU held the entire model. Each parallelism scheme breaks that in a different way.
| Layout | What each rank holds | How to get the global norm |
|---|---|---|
| DDP | Full, already-averaged gradients | Local norm is already global |
| FSDP / ZeRO-2/3 | A shard of every gradient | All-reduce the sum of squares, then sqrt |
| Tensor parallel | Shards of split weights, copies of replicated ones (norms, biases) | Sum squares of shards once, count replicated params on one rank only |
| Pipeline parallel | Whole layers for one stage | All-reduce the sum of squares across stages |
FSDP1 exposes FullyShardedDataParallel.clip_grad_norm_(), which does the cross-rank reduction; calling the plain utility on an FSDP1 module computes a per-shard norm and clips each rank differently. Under FSDP2 the parameters are DTensors, and the plain torch.nn.utils.clip_grad_norm_ reduces across the sharding mesh and returns a DTensor norm. Pipeline stages still need an explicit reduction. Megatron-style frameworks do this inside their optimizer wrapper. If you write your own, the pattern is short:
def global_grad_norm(params, groups, replicated):
"""params: local grads; groups: process groups to reduce over (dp shard, tp, pp);
replicated: set of params that are copies on every TP rank (count once)."""
sq = torch.zeros((), device="cuda", dtype=torch.float32)
for prm in params:
if prm.grad is None:
continue
if prm in replicated and tp_rank() != 0:
continue # avoid counting copies tp_size times
sq += prm.grad.detach().float().pow(2).sum()
for g in groups:
dist.all_reduce(sq, op=dist.ReduceOp.SUM, group=g)
return sq.sqrt()
norm = global_grad_norm(local_params, [dp_shard_group, tp_group, pp_group], replicated)
coef = torch.clamp(max_norm / (norm + 1e-6), max=1.0)
for prm in local_params:
if prm.grad is not None:
prm.grad.mul_(coef) # same coef on every rankAccumulate in fp32 even when gradients are bf16, or the sum of millions of squares loses precision. And verify the result: run a small model both unsharded and sharded on the same batch and compare the norms. They should agree to a few digits. A sharded norm too small by roughly the square root of the TP degree means the shards' squares were never summed across the TP group. A norm only slightly too large usually means replicated parameters were counted on every rank, since they are a small share of the total.
Clipping and Adam
With plain SGD, clipping caps the step length directly. With Adam the story is subtler, because Adam divides the first moment by the square root of the second moment. Scaling a single gradient by a constant does not cancel out, though, because both moments are exponential averages over many steps.
Consider a spike step where the raw norm is 50 times normal. Unclipped, the first moment jumps and the update in that step is large. Worse, the squared spike lands in the second moment, which has a long memory (with beta2 = 0.95 or 0.999, roughly 20 or 1,000 steps). For that whole window the denominator is inflated, so the effective learning rate on the affected parameters collapses. The run looks stable but learns slowly afterwards. Clipping caps what enters both moments, so it protects the optimizer state as well as the current step. That is the main reason clipping still matters in adaptive optimizers.
It also explains why the threshold should sit just above the typical norm, not far above it. A threshold at 100 times normal never fires, and the moments absorb every spike.
Choosing the threshold
A fixed threshold of 1.0 is the most common choice in published large language model recipes; GPT-3 and Llama 2 both report clipping at 1.0. Those recipes use per-token mean losses, and the right value for your run depends on its own norm scale, so 1.0 is a starting point, not a law. The more reliable method is to measure: run a few hundred steps without clipping (or with a very loose threshold), record the norm distribution after warm-up, and set the threshold near a high percentile so it fires on a few percent of steps.
Adaptive schemes avoid the absolute scale. A rolling-statistics clip compares each step's norm to an exponential moving average and clips at a multiple of it, which follows the natural drift of the norm over training. Adaptive gradient clipping (AGC, introduced with NFNets by Brock et al., 2021) clips each unit's gradient relative to the norm of that unit's weights, which made batch-norm-free image models trainable at large batch sizes. A rolling clip is a few lines:
class RollingClip:
def __init__(self, k=3.0, beta=0.99, warmup=200, floor=0.05):
self.k, self.beta, self.warmup, self.floor = k, beta, warmup, floor
self.ema, self.n = None, 0
def threshold(self):
if self.ema is None or self.n < self.warmup:
return 1.0 # fixed clip while statistics settle
return max(self.floor, self.k * self.ema)
def update(self, norm):
norm = min(norm, self.threshold()) # do not let spikes inflate the average
self.ema = norm if self.ema is None else self.beta * self.ema + (1 - self.beta) * norm
self.n += 1Feed the clipped value into the average, as above, or a run of spikes raises the threshold until clipping stops protecting anything.
Worked example: one spike, then a drift
Take a 1B-parameter model trained with AdamW, max_norm 1.0, and a logged pre-clip norm. After warm-up the norm sits around 0.35 with a 99th percentile near 0.8, so clipping fires on well under 1 percent of steps. At step 41,200 the norm reads 9.6. The coefficient is 1.0 / 9.6, about 0.104, so every gradient is multiplied by roughly a tenth and the update keeps its direction at a tenth of its length. The loss ticks up by 0.03 and recovers within a hundred steps. Without clipping, the second moment would have absorbed a squared spike about 750 times its usual size (9.6 squared over 0.35 squared), suppressing updates for hundreds of steps.
Now suppose the clip fraction climbs from 1 percent to 40 percent over a day, with the median norm creeping up to 1.3. Clipping is no longer catching rare spikes; it is silently acting as a learning-rate reduction on most steps. That is a symptom, usually of a learning rate too high for the current phase, a data shard with corrupt or extremely long samples, or growing attention logits. The fix is upstream; raising the threshold only hides it.
Clipping as a health signal
Log three numbers every step: the pre-clip norm, whether clipping fired, and the coefficient. From them derive the clip fraction over a window. Optionally log per-layer norms every few hundred steps; a spike that starts in the embedding or final layer norm points at data, while one in the deepest attention blocks often points at logit growth. When the norm spikes, record the batch indices so the offending samples can be inspected and, if needed, skipped on a restart. Some large-run recipes go further and skip the update entirely when the norm exceeds a large multiple of its running average, rather than merely clipping it.
Failure modes
- Clipping scaled fp16 gradients. The clip fires every step and training crawls. Call
unscale_first. - Per-shard norms under FSDP1 or ZeRO. Each rank clips with a different coefficient and replicas drift apart. Use the wrapper's method or reduce the sum of squares yourself.
- Replicated parameters counted per TP rank. The norm is inflated, so the effective threshold is too tight.
- Clipping per micro-batch. Accumulated partial sums are capped separately and the total can exceed the threshold.
- Logging the post-clip norm. Spikes disappear from dashboards.
- NaN propagation. Without a scaler, a non-finite norm turns every gradient into NaN; check and skip.
- Threshold copied across loss conventions. Switching from mean to summed loss multiplies the norm by the token count, and the old threshold clips everything.
Trade-offs
Clipping costs one extra pass over the gradients plus, when sharded, one scalar all-reduce per parallel group. On large models the pass is a few milliseconds and the all-reduce is latency-bound, so the overhead is small but not zero; fused or foreach implementations keep it down. A tight threshold makes training more robust but slows it, because every clipped step is a smaller step. A loose one lets the moments absorb spikes. Adaptive clipping removes tuning but adds state that must be checkpointed with the optimizer, or a restart resets it. And clipping never fixes a cause; it buys time to find it. For the wider debugging workflow see debugging GPU training runs, and for how the underlying update rule behaves see gradient descent.
What to do next
- Find your clip call and confirm the order: accumulate, sync,
unscale_, clip, step. See mixed-precision training for the scaler. - Log the pre-clip norm, a clipped flag and the coefficient every step.
- Measure the norm distribution after warm-up and set the threshold near the 97th to 99th percentile.
- Under sharding, check the global norm against an unsharded run on one batch. For FSDP details see FSDP in depth.
- Alert when the windowed clip fraction exceeds about 5 percent, and treat it as a bug report.
- Add a non-finite check that skips the step and records the batch.