skip to content

Where does torch.nn.utils.clip_grad_norm_ go in a PyTorch loop, and what does it clip?

level: seniorimportance: should knowfreq 44%

answer

  1. between backward and step
  2. one norm over everything, one factor
  3. direction kept, length capped
  4. the return value is worth logging
  5. unscale first when using float16 AMP

basics

~20 s

Call it after loss.backward() and before optimizer.step(). It computes one global L2 norm across all the gradients you pass it and, if that norm exceeds max_norm, scales every gradient by the same factor so the combined norm equals max_norm.

solid answer

~50 s

`torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)` edits `.grad` in place, so it must sit between `loss.backward()`, which creates the gradients, and `optimizer.step()`, which consumes them. It treats all the listed parameters' gradients as one long vector, takes its L2 norm, and if that norm exceeds `max_norm` multiplies **every** gradient by `max_norm / total_norm`. Because a single factor is applied, the direction of the update is preserved and only its length is capped — which is what makes it safe. The function returns the total norm **before** clipping, and logging that value is the cheapest instability detector you have: a healthy run shows a stable band with occasional spikes, and a run about to diverge shows the norm climbing steadily. `clip_grad_value_` is the blunter alternative that clamps each element independently and does change the update direction. Under `torch.amp` float16 you must call `scaler.unscale_(optimizer)` first, or you are clipping loss-scaled gradients and `max_norm` means nothing.

code

python · 12 lines
python
import torch
from torch import nn

model = nn.Linear(4, 2)
opt = torch.optim.SGD(model.parameters(), lr=0.1)

opt.zero_grad(set_to_none=True)
loss = model(torch.randn(8, 4)).pow(2).mean()
loss.backward()
total = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
print("norm before clipping:", total.item())
opt.step()

go deeper

for a junior

Know that clipping exists to stop one bad batch from blowing up the weights, and that the call goes after backward and before the optimizer step.

for a middle

Explain the global-norm semantics — one norm, one scaling factor, direction preserved — and contrast it with per-element clipping by value.

for a senior

Use the returned norm as a diagnostic, set the threshold from measurement rather than folklore, and get the unscale-before-clip ordering right under mixed precision.

for a principal

Position clipping honestly in a stability strategy: it is a shock absorber for tail batches, and a run that clips constantly is telling you the learning rate, warmup or initialization is wrong at a level clipping cannot repair.

## Where it goes Gradient clipping is an in-place edit of the `.grad` tensors. That fixes its position in the loop exactly: after `loss.backward()` has populated them and before `optimizer.step()` reads them. Put it before backward and there is nothing to clip; put it after the step and you have modified gradients that were already used, which does nothing except confuse the next `zero_grad`. The common signature is `torch.nn.utils.clip_grad_norm_(parameters, max_norm, norm_type=2.0)`. The trailing underscore is PyTorch's in-place convention. `parameters` can be any iterable of tensors — usually `model.parameters()`, but for a run with parameter groups you may want to clip each group separately. ## Global norm, not per-tensor The key semantic is that this is a **global** clip. All the gradients are conceptually concatenated into one vector; its norm is computed once; and if that norm is above the threshold, one scalar factor `max_norm / total_norm` scales every gradient tensor. Nothing is clipped when the norm is already under the threshold. That design is deliberate. Because a single factor is applied everywhere, the *direction* of the parameter update is exactly preserved — you take a shorter step along the same line, which is a much gentler intervention than distorting the direction. This is why norm clipping is the default for transformer training while per-element clipping is rare. `clip_grad_value_(parameters, clip_value)` is the alternative: it clamps every gradient element into `[-clip_value, clip_value]` independently. It is simple and it bounds the worst element, but it changes the direction of the update whenever it fires, and it is much harder to reason about because the effective step depends on how many elements were clamped. ## Reading the returned norm `clip_grad_norm_` returns the total norm computed **before** any scaling, as a tensor. This return value is genuinely useful and routinely ignored. Log it every step (or every N steps) alongside the loss: - A stable band with occasional spikes is healthy — the spikes are the batches clipping was added for. - A norm that grows monotonically over hundreds of steps is a run heading for divergence; clipping is masking the symptom and the learning rate or the initialization is the real problem. - A norm that is *always* far above `max_norm` means you are effectively training with a normalized-gradient method and a much smaller effective learning rate than you think. - A norm that suddenly becomes `nan` or `inf` pinpoints the step where numerics broke, which is far more actionable than noticing the loss went NaN twenty steps later. The `error_if_nonfinite` argument makes that failure loud instead of silent. ## Interaction with mixed precision Under `torch.amp` with float16, `scaler.scale(loss).backward()` leaves `.grad` multiplied by the loss scale — a factor that can be tens of thousands and that changes over time. Clipping those directly compares a scaled norm against your `max_norm`, so the clip fires on essentially every step and the threshold is meaningless. The correct order is: 1. `scaler.scale(loss).backward()` 2. `scaler.unscale_(optimizer)` — divides `.grad` by the current scale 3. `torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm)` 4. `scaler.step(optimizer)` — sees the gradients are already unscaled and does not repeat it 5. `scaler.update()` This is one of the highest-yield details on the whole leaf, because the bug is invisible: the model still trains, just badly. ## Choosing max_norm There is no universal value, but there is a reliable method: run a few hundred steps without clipping, log the returned norm, and set `max_norm` somewhere above the typical value so that clipping fires only on the tail. Setting it far below the typical norm turns every step into a normalized-gradient step; setting it far above means it never fires and buys nothing. `1.0` is the conventional starting point for transformer language-model training largely because the recipes that popularized it used it, not because it is derived from anything. ## What clipping cannot fix Clipping bounds the *size* of an update; it does not make an unstable objective stable. If the loss is NaN because of a log of zero, a division by a near-zero variance, or float16 overflow inside the forward pass, the gradient is already `nan` before clipping sees it — and scaling `nan` by any factor is still `nan`. Clipping is a shock absorber for occasional outlier batches, not a cure for a broken model or a learning rate that is fundamentally too high.

  • How does clip_grad_value_ differ from clip_grad_norm_?
    clip_grad_value_ clamps every gradient element independently into [-clip_value, clip_value]. It bounds the worst element but changes the direction of the update whenever it fires, because different coordinates are scaled by different amounts. clip_grad_norm_ applies one factor to everything, preserving direction and only shortening the step, which is why it is the default in most modern recipes.
  • How do you choose max_norm for a new model?
    Measure it. Train a few hundred steps with clipping effectively off, log the norm that clip_grad_norm_ returns, and set the threshold above the typical value so it fires only on the tail of the distribution. If it fires on nearly every step, the threshold is too tight and you have silently turned the run into normalized-gradient descent with a much smaller effective step size.
  • Does gradient clipping prevent NaN losses?
    No. If a NaN or inf entered during the forward or backward pass — a log of zero, a zero-variance normalization, a float16 overflow — the gradient is already non-finite, and scaling a NaN leaves a NaN. Pass error_if_nonfinite=True to make that case raise at the exact step it appears rather than surfacing as a NaN loss much later.

It is a speed limiter, not a steering correction: the car keeps heading exactly where it was pointed, just no faster than the cap.

saying these in an interview costs you the question

  • Calling clip_grad_norm_ before loss.backward()
  • Thinking each parameter tensor is clipped to its own norm
  • Believing clipping fixes NaN gradients
  • Clipping AMP gradients without calling scaler.unscale_ first
  • Ignoring the returned norm and never measuring the threshold

context