Debugging training
Read loss curves, run the sanity checks experts use, and fix exploding gradients and NaNs.
- Diagnose underfitting, overfitting and divergence from loss curves
- Use sanity checks: the initial loss, overfitting one batch, and inspecting data
- Clip gradients by global norm and track down NaNs
Neural networks fail silently. A bug rarely crashes; it just makes the network learn a bit worse, and you can lose days. Experienced practitioners debug methodically, starting with the loss curves:
| What you see | Likely cause | Try |
|---|---|---|
| Loss explodes or becomes NaN | learning rate too high, bad data, log(0) | lower the LR, clip gradients, check inputs |
| Loss flat from the start | learning rate too low, bug in backward pass, dead ReLUs | raise the LR, gradient check, check init |
| Train and validation both high | underfitting: model too small or not trained enough | bigger model, train longer |
| Train low, validation rising | overfitting | regularize, more data, early stopping |
| Very noisy loss | batch too small, LR too high | bigger batch, lower LR |
Try it
Diagnose the training run
Each description is what a keeper at Doodle Lab saw on their loss chart. What’s the most likely problem?
“Loss: 2.3, 2.1, 9.8, 140, NaN”
“Train 0.02, validation 1.40 and rising”
“Train 1.9, validation 1.9, barely moving after 50 epochs, tiny model”
“A 10-class classifier starts at loss 0.0001 on its very first batch”
“Validation accuracy is exactly 10% on a 10-class problem forever”
“Loss bounces between 0.8 and 3.0 every few steps and never settles”
Sanity checks that save days
- Check the initial loss. With K classes and random weights the predictions are near uniform, so the cross-entropy should start near (2.30 for 10 classes). Far off? Something’s wrong before training even begins.
- Overfit one batch. Train on a single batch of a few examples. A working network drives the loss to almost zero. If it can’t, there’s a bug - no point training on everything.
- Look at the data. Plot a few inputs with their labels, after all preprocessing. Mislabeled data, wrong normalization and accidentally shuffled labels are depressingly common.
- Start simple. A small model and a known-good learning rate first; add complexity once it works.
1import numpy as np
2
3rng = np.random.default_rng(0)
4for classes in (2, 10, 1000):
5 logits = rng.normal(0, 0.01, size=(256, classes)) # small random init
6 labels = rng.integers(0, classes, size=256)
7 shifted = logits - logits.max(axis=1, keepdims=True)
8 log_probs = shifted - np.log(np.exp(shifted).sum(axis=1, keepdims=True))
9 loss = -log_probs[np.arange(256), labels].mean()
10 print(f"{classes:>4} classes: initial loss {loss:.3f}, ln(K) = {np.log(classes):.3f}")2 classes: initial loss 0.693, ln(K) = 0.693 10 classes: initial loss 2.303, ln(K) = 2.303 1000 classes: initial loss 6.907, ln(K) = 6.908
Exploding gradients and gradient clipping
Sometimes a single unlucky batch produces a gigantic gradient, and one step undoes hours of training. Gradient clipping caps the size of the update: compute the global norm of all gradients together, and if it exceeds a threshold, scale every gradient down by the same factor. The direction is preserved; only the length is limited. It’s standard for recurrent networks and transformers.
1import numpy as np
2
3def clip_by_global_norm(grads, max_norm):
4 total = np.sqrt(sum(np.sum(g ** 2) for g in grads))
5 scale = min(1.0, max_norm / (total + 1e-12))
6 return [g * scale for g in grads], total
7
8grads = [np.array([3.0, 4.0]), np.array([[12.0]])] # global norm = 13
9clipped, norm = clip_by_global_norm(grads, max_norm=1.0)
10print(f"norm before: {norm:.1f}")
11print([np.round(g, 3).tolist() for g in clipped])
12print(f"norm after: {np.sqrt(sum(np.sum(g ** 2) for g in clipped)):.1f}")norm before: 13.0 [[0.231, 0.308], [[0.923]]] norm after: 1.0
Key takeaways
Read the loss curves first: divergence, flat lines, underfitting and overfitting each look different.
The initial loss should be about ln(K); a network should be able to overfit one small batch.
Look at your data after preprocessing - many “model bugs” are data bugs.
Clip gradients by global norm to stop single bad steps; hunt NaNs at the first non-finite loss.
Lesson quiz
7 questions · pass with 5 correct · up to 50 XP
Passing this quiz completes the lesson and keeps your streak going. Questions you miss come back in review sessions later.
Practice: write Python
Write Python in the editor and run it against sample inputs. Python runs locally in your browser using a WebAssembly runtime.
Clip by global norm
The first input line is max_norm; each following line is one parameter’s gradient (a list of numbers). Compute the global norm of all gradients together and, if it exceeds max_norm, scale every gradient by max_norm / norm. Print the norm before (3 decimals), each clipped gradient to 3 decimals, and the norm after.
- Too big
- Already small
Python runs in a sandboxed browser worker with a 60 second time limit. Its runtime loads from the Pyodide CDN; your code stays in this browser.
Diagnose the run
The input has two lines: training losses per epoch, then validation losses. Diagnose with these rules, checked in order, and print the verdict:
- any loss is
nanorinf, or the last training loss is above the first →diverging: lower the learning rate - the last training loss is more than 90% of the first →
not learning: check the learning rate and the backward pass - the last validation loss is more than 10% above the best validation loss →
overfitting: regularize or stop early at epoch N(N = the best validation epoch, counting from 1) - otherwise →
healthy
- Overfitting
- Diverging
- Stuck
- Healthy
Python runs in a sandboxed browser worker with a 60 second time limit. Its runtime loads from the Pyodide CDN; your code stays in this browser.
Questions about this lesson
Stuck? Ask. Figured something out? Share it. Explaining is one of the best ways to learn.
Loading posts…