Debugging PyTorch

Tracking down a NaN loss

NaN spreads through every calculation it touches, so the job is finding the first infinity — with a finiteness guard, anomaly detection, and a short list of usual suspects.

On this page 5
  1. Why it exists
  2. How it works
  3. Where you have seen this
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

NaN — "not a number" — is what a calculation produces when asked something impossible, and one NaN poisons every calculation it touches afterwards.

It behaves like a drop of spoiled milk in the pot. The moment it lands, everything it mixes with is spoiled too, and no later stirring un-spoils it. By the time you smell it, the whole pot is gone.

So the debugging question is never "why is my loss NaN now?" It is: where did the very first drop land?

Why it exists

Computers store numbers in boxes of fixed size. Sometimes you ask for something no box can hold. A result grown beyond the largest storable value, say, or a genuinely undefined operation like dividing zero by zero. The machine then records a special marker instead of a number.

There are two markers. inf means "grew past the largest box" — infinity. NaN means "no defined answer exists". They convert into each other easily: subtract infinity from infinity, or multiply infinity by zero, and NaN is born. In training, the story is usually: something overflowed to inf first, and a step later the infs collided into NaN.

How it works

step 1: loss 900        fine
step 2: loss 196,761    growing... (the real bug is already here)
step 3: loss inf        overflowed
step 4: loss nan        poisoned — and every weight now nan too

Notice the bug predates the NaN by several steps. The NaN is the smoke; the fire started earlier.

Where you have seen this

Calculators show Error when asked to divide by zero — same marker, friendlier costume. And once a spreadsheet cell shows an error, every formula referencing it errors too, cascading down the sheet exactly like NaN cascades through training.

Remember this

  • NaN spreads: one bad value contaminates everything downstream.
  • It is usually born as inf — an overflow — one or more steps earlier.
  • Hunt the first non-finite value, not the step where you noticed.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Captured with torch 2.5.1 on CPU; seeded and deterministic.

Watching a loss explode in real time

Unscaled inputs and a hot learning rate — the most common NaN recipe in the wild:

loss_explosion.py
import torch
import torch.nn as nn

torch.manual_seed(0)
X = torch.randn(256, 10) * 10                # unscaled inputs: ten times too big
y = X.sum(dim=1, keepdim=True)

model = nn.Sequential(nn.Linear(10, 32), nn.ReLU(), nn.Linear(32, 1))
optimizer = torch.optim.SGD(model.parameters(), lr=0.1)   # too hot for this data
loss_fn = nn.MSELoss()

for step in range(8):
    optimizer.zero_grad()
    loss = loss_fn(model(X), y)
    loss.backward()
    optimizer.step()
    print(f"step {step}  loss {loss.item():.1f}")
    if not torch.isfinite(loss):
        print("loss is no longer a number; everything after this is garbage")
        break
Output
step 0  loss 900.1
step 1  loss 196761.6
step 2  loss 154129986658238464.0
step 3  loss inf
loss is no longer a number; everything after this is garbage

The signature to learn: loss growing by orders of magnitude per step, then inf. The torch.isfinite guard catches both inf and NaN, and stopping at first detection preserves the crime scene — one step later, the weights themselves are contaminated and the evidence is gone.

When the culprit hides: anomaly detection

Explosions announce themselves. The sneakier NaNs come from one impossible operation inside an otherwise calm run — a log of zero, a division by zero, a square root of a negative. For those, PyTorch has a bloodhound:

anomaly_hunt.py
import torch

torch.autograd.set_detect_anomaly(True)      # slow: switch on only while hunting

torch.manual_seed(0)
emb = torch.randn(3, 4, requires_grad=True)
emb2 = emb * torch.tensor([[1.], [1.], [0.]])   # row 2 is a padding row: all zeros
normed = emb2 / emb2.norm(dim=1, keepdim=True)  # 0 / 0 = nan for the padding row
loss = normed.sum()
loss.backward()
Output
UserWarning: Error detected in DivBackward0. Traceback of forward call that caused the error:
  File "anomaly_hunt.py", line 8, in <module>
    normed = emb2 / emb2.norm(dim=1, keepdim=True)  # 0 / 0 = nan for the padding row
...
RuntimeError: Function 'DivBackward0' returned nan values in its 1th output.

(Traceback trimmed; the file paths and line numbers will be your own.) Two gifts in that output: the operation type (DivBackward0 — a division) and, above it, the exact source line of the forward pass that created the poison. Normalising a zero vector — a padded row, an empty segment — is a genuinely common real-world case of this bug.

The usual suspects

CauseSignatureFix
Learning rate too highLoss grows, then infLower lr; run the range test
log/sqrt of zero, division by zeroCalm run, sudden NaNAdd an epsilon: x.clamp_min(1e-8)
Unscaled inputs or targetsHuge loss from step 0Normalise features and targets
Exploding gradients (deep/recurrent nets)Intermittent spikes, then infGradient clipping
float16 training without a scalerNaN under autocastUse GradScaler, or bf16

Common mistakes

Restarting with a lower learning rate and no diagnosis. It often works — and teaches nothing, so the NaN returns on the next dataset. Spend the five minutes finding the actual suspect from the table.

Leaving set_detect_anomaly(True) on permanently. It roughly doubles-to-quadruples step time and is meant as a temporary bloodhound, not a smoke detector. On for the hunt, off for training.

Guarding with loss.item() != loss.item()-style NaN checks only. That misses inf, which arrives first. torch.isfinite(loss) catches the whole family.

Clamping the loss itself. loss.clamp(max=1e6) hides the symptom and destroys the gradient signal. Clamp inputs to fragile ops (logs, divisions, square roots), never the loss.

Try it yourself

Fix the explosion two separate ways and confirm each works alone: divide X by 10 (scale the data), or set lr=0.001 (cool the optimiser). Then fix the anomaly demo with emb2.norm(dim=1, keepdim=True).clamp_min(1e-8) and verify loss.backward() runs clean.

What to learn next

Researcher — Mathematics and papers.

IEEE 754 semantics

Float formats reserve exponent patterns for $\pm\infty$ and NaN. Generation rules: overflow past the format maximum produces $\pm\infty$; the indeterminate forms $0/0$, $\infty - \infty$, $0 \times \infty$, $\sqrt{x<0}$, $\log(x<0)$ produce NaN. Propagation: any arithmetic with NaN yields NaN, and NaN compares unequal to everything including itself (the basis of the x != x test). One consequential subtlety: max(NaN, x) semantics differ across libraries — reductions may or may not swallow NaNs — so torch.isfinite on the tensor, not comparisons, is the reliable audit.

Format ranges explain the fp16 story: float16 maxes at 65,504, so attention scores or losses overflow with ease; bfloat16 keeps float32's exponent (max $\approx 3.4 \times 10^{38}$) at the cost of mantissa, which is why bf16 training needs no loss scaling while fp16 does — Micikevicius et al. (2018), Mixed Precision Training, is the loss-scaling reference, implemented as torch.amp.GradScaler.

Exploding gradients

For recurrent or very deep compositions, the gradient norm is bounded by products of per-layer Jacobian norms; spectral radius above 1 compounds geometrically — Pascanu et al. (2013), On the difficulty of training recurrent neural networks, formalise this and propose norm clipping: rescale $g \leftarrow g \cdot \tau / |g|$ when $|g| > \tau$. Clipping bounds the update, making the divergence loop (big step → worse curvature region → bigger gradient) self-limiting; it does not repair an already-NaN gradient, hence clip and fix the source.

Softmax and log-sum-exp stability

The canonical hidden NaN factory: $\mathrm{softmax}(z)$ computed naively overflows $e^z$ near $z \approx 89$ in float32; the max-subtraction identity keeps exponents non-positive. Fused implementations (CrossEntropyLoss, log_softmax, flash-attention's online softmax) exist substantially to guarantee this — the reason the logits lesson insists losses receive raw logits.

Hunting machinery

torch.autograd.set_detect_anomaly(True) wraps each backward node with a finiteness check and retains forward stack traces — O(graph) overhead, the measured 2–4x slowdown. Cheaper standing guards: torch.isfinite(loss) per step (negligible cost, catches the event within one step), periodic parameter audits any(not p.isfinite().all() for p in model.parameters()), and forward hooks asserting finiteness per module for localisation without anomaly mode — the hook pattern from forward and backward hooks.

What to learn next