Error database

Loss is NaN during training

One arithmetic step produced not-a-number and it spread to every weight. Lower the learning rate first, then check your data and your loss function.

The message you saw
Loss is NaN during training

By Updated

The error

There is no exception. The training loop keeps running, which is what makes this one unpleasant.

Output
epoch 1 | step  100 | loss 2.3014
epoch 1 | step  200 | loss 1.8872
epoch 1 | step  300 | loss nan
epoch 1 | step  400 | loss nan
epoch 1 | step  500 | loss nan

Sometimes it arrives via infinity first, and sometimes a library says it out loud:

Output
epoch 1 | step  300 | loss inf
epoch 1 | step  400 | loss nan
Output
ValueError: Input contains NaN.

What it means

Somewhere in the forward or backward pass, an operation produced NaN — "not a number", the value floating point arithmetic returns for undefined results like 0 divided by 0, infinity minus infinity, or the logarithm of zero.

NaN is contagious. Anything computed from a NaN becomes NaN. So one bad value in one step becomes a NaN gradient, which becomes NaN weights after optimizer.step(), and from that moment every output of the model is NaN. The model is dead — no amount of further training recovers it, and you must restart from a checkpoint or from scratch.

The step where you see NaN is not the step where it started. It is the step where it finished spreading.

Why it happens

The learning rate is too high. This causes more NaN losses than everything else combined. A step that overshoots the minimum lands somewhere with a larger error, whose gradient is larger, whose next step overshoots further. Within a few iterations the weights reach infinity, and infinity minus infinity is NaN. If your loss climbed for a few steps before going bad, this is almost certainly what happened.

A logarithm of zero. Cross-entropy takes the log of the predicted probability. If a probability reaches exactly 0, the log is negative infinity. This is why applying softmax yourself and then passing the result to nn.CrossEntropyLoss is a bug: that function applies log-softmax internally, in a numerically stable way, and a second softmax destroys that protection. The same applies to sigmoid followed by nn.BCELoss instead of nn.BCEWithLogitsLoss.

NaN was already in the data. Missing values in a CSV, a division during feature engineering, or normalising with a standard deviation of zero. The model faithfully propagates it.

Division by zero or a square root of zero in a custom loss. torch.sqrt(0) returns 0, but its gradient at zero is infinite, so this one fails on the backward pass while the forward pass looks fine — a genuinely confusing failure mode.

Float16 overflow. In mixed precision, values above roughly 65,504 become infinity. This is exactly what GradScaler exists to manage, and skipping it while using float16 leads here.

Exploding gradients in recurrent models. Gradients multiplied through many time steps grow without limit. RNNs and LSTMs on long sequences are the classic case.

How to fix it

1. Lower the learning rate by a factor of 10, then run again. Free, fast, and correct most of the time.

python
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)   # was 1e-3

If 1e-4 also produces NaN, try 1e-5. If a tiny learning rate still fails, the cause is elsewhere and you should move to step 2.

2. Check your data before blaming the model.

python
for x, y in train_loader:
    assert torch.isfinite(x).all(), "input contains NaN or inf"
    assert torch.isfinite(y).all(), "target contains NaN or inf"
    break

print(x.min().item(), x.max().item(), x.mean().item())

Inputs in the hundreds or thousands (raw pixel values 0-255, unscaled prices) make the first gradients enormous. Normalise them:

python
x = (x - x.mean()) / (x.std() + 1e-8)     # the epsilon guards a constant feature

For pandas data, df.isna().sum() and np.isinf(df.select_dtypes("number")).sum() locate the columns.

3. Pass raw logits to the loss, not probabilities. This is a real bug, not a style preference.

python
# wrong — two softmaxes, and log(0) becomes possible
logits = model(x)
probs = torch.softmax(logits, dim=1)
loss = nn.CrossEntropyLoss()(probs, y)

# correct — CrossEntropyLoss applies log-softmax internally and stably
logits = model(x)
loss = nn.CrossEntropyLoss()(logits, y)

The binary equivalent: use nn.BCEWithLogitsLoss() on raw outputs rather than nn.BCELoss() on sigmoid outputs. Remember to apply softmax or sigmoid yourself at prediction time, when you want probabilities to show a user.

4. Clip the gradients. This puts a ceiling on the size of any update, which stops the runaway loop directly. It is standard practice for RNNs and transformers.

python
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()

The order matters: after backward(), before step(). With mixed precision you must unscale first, otherwise you are clipping scaled values:

python
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()

5. Find the exact step where NaN first appears. Add a check inside the loop rather than reading the whole log afterwards:

python
for step, (x, y) in enumerate(train_loader):
    loss = criterion(model(x), y)
    if not torch.isfinite(loss):
        print(f"loss became {loss.item()} at step {step}")
        print("input finite:", torch.isfinite(x).all().item())
        print("target range:", y.min().item(), y.max().item())
        torch.save({"x": x, "y": y}, "bad_batch.pt")   # keep the batch to inspect
        break
    loss.backward()

Saving the offending batch is worth the two lines. Being able to reproduce it on demand turns an overnight mystery into a five-minute inspection.

6. If the forward pass is clean but the gradients are not, use anomaly detection. It points at the operation that produced the NaN in the backward pass.

python
torch.autograd.set_detect_anomaly(True)      # debugging only — several times slower

Turn it off again once you have the answer.

7. Guard the unstable operations in custom losses.

python
loss = torch.sqrt(diff.pow(2).sum() + 1e-8)      # avoids an infinite gradient at zero
ratio = numerator / (denominator + 1e-8)
logp  = torch.log(p.clamp(min=1e-8))

8. If you are using float16 mixed precision, switch to bfloat16 where the hardware allows it. bfloat16 has the same numeric range as float32, so overflow largely stops being a concern and no loss scaler is needed. It is available on Ampere (RTX 30-series, A100) and newer.

python
with torch.amp.autocast("cuda", dtype=torch.bfloat16):
    loss = criterion(model(x), y)

9. For transformers, add a learning-rate warmup. Starting from near zero and rising over the first few hundred steps avoids the large early updates that destabilise attention layers. Every standard transformer recipe includes this for a reason.

How to prevent it

Log the loss every 50 steps and look at the first few hundred. A loss that rises before it breaks is a learning-rate problem; a loss that turns to NaN from a good value in one step is a data or numerics problem. Knowing which halves your search.

Before a long run, overfit a single batch on purpose. Take one batch, train on it for 200 steps, and confirm the loss drops close to zero. This takes a minute and catches a large share of loss-function bugs, wrong label ranges and unstable custom operations before you commit hours of GPU time.

Then keep three defaults in place: normalised inputs, gradient clipping at 1.0, and checkpoints saved every epoch so a NaN at hour six costs you one epoch instead of six hours. A torch.isfinite(loss) check with a break is four lines that stop a dead model from burning a whole night of compute.