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.
- 8 min read
- 3 reading levels
- Published
Read these first
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 tooNotice 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
- Gradient clipping — the standing defence against the explosion pathway.
- Checking that gradients reach every layer — the opposite disease: gradients that die instead of exploding.
- Finding a learning rate that works — staying on the safe side of the cliff by measurement.
Developer — Code and libraries.
Setup
pip install torchCaptured 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:
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")
breakstep 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:
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()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
| Cause | Signature | Fix |
|---|---|---|
| Learning rate too high | Loss grows, then inf | Lower lr; run the range test |
log/sqrt of zero, division by zero | Calm run, sudden NaN | Add an epsilon: x.clamp_min(1e-8) |
| Unscaled inputs or targets | Huge loss from step 0 | Normalise features and targets |
| Exploding gradients (deep/recurrent nets) | Intermittent spikes, then inf | Gradient clipping |
| float16 training without a scaler | NaN under autocast | Use 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
- Gradient clipping — the standing defence against the explosion pathway.
- Checking that gradients reach every layer — the opposite disease: gradients that die instead of exploding.
- Finding a learning rate that works — staying on the safe side of the cliff by measurement.
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
- Gradient clipping — the standing defence against the explosion pathway.
- Checking that gradients reach every layer — the opposite disease: gradients that die instead of exploding.
- Finding a learning rate that works — staying on the safe side of the cliff by measurement.