Um momento
0xB0Lesson 12 of 17

Debugging training

Read loss curves, run the sanity checks experts use, and fix exploding gradients and NaNs.

26 min 7-question quiz 2 code exercises
By the end of this lesson you can
  • 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 seeLikely causeTry
Loss explodes or becomes NaNlearning rate too high, bad data, log(0)lower the LR, clip gradients, check inputs
Loss flat from the startlearning rate too low, bug in backward pass, dead ReLUsraise the LR, gradient check, check init
Train and validation both highunderfitting: model too small or not trained enoughbigger model, train longer
Train low, validation risingoverfittingregularize, more data, early stopping
Very noisy lossbatch too small, LR too highbigger 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?

0 of 6 sortedScore 0/0
  • “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

  1. Check the initial loss. With K classes and random weights the predictions are near uniform, so the cross-entropy should start near ln⁡K\ln K (2.30 for 10 classes). Far off? Something’s wrong before training even begins.
  2. 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.
  3. 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.
  4. Start simple. A small model and a known-good learning rate first; add complexity once it works.
initial_loss.py
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}")
Output
   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.

clip.py
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}")
Output
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.

Exercise 1

Clip by global norm

+25 XP

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
main.py
Loading editor…

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.

Exercise 2

Diagnose the run

+25 XP

The input has two lines: training losses per epoch, then validation losses. Diagnose with these rules, checked in order, and print the verdict:

  1. any loss is nan or inf, or the last training loss is above the first → diverging: lower the learning rate
  2. the last training loss is more than 90% of the first → not learning: check the learning rate and the backward pass
  3. 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)
  4. otherwise → healthy
  • Overfitting
  • Diverging
  • Stuck
  • Healthy
main.py
Loading editor…

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…

Gostou da aula? 😆👍
Apoie nosso trabalho com uma doação: