Training Stability Techniques

Reviewed & published by Brayan K

By the end of this lesson you'll be able to spot an unstable training run, clip gradients to a safe norm, warm up and schedule the learning rate, and make every run reproducible — so deep learning actually converges instead of blowing up.

Part of the free AI & Machine Learning course at LearnCodingFast — hands-on lessons with examples you run in your browser, plus practice exercises and a quick quiz.

What You'll Learn in This Lesson

🏗️ Real-World Analogy: keeping a tall tower from toppling

Think of a deep network as a very tall tower of building blocks, one block per layer. Training nudges the blocks to make the tower taller and straighter. The taller it gets, the easier it is to topple — and the same techniques a builder uses map exactly onto stable training:

1 Exploding and Vanishing Gradients

During backpropagation the gradient is multiplied through every layer. Exploding gradients happen when those factors are larger than 1 on average — the product grows huge, the weights jump, and the loss blows up to inf then NaN. Vanishing gradients are the opposite: factors below 1 shrink the product toward zero, so the early layers barely update and the network stops learning.

You spot exploding gradients by watching the gradient norm climb and the loss spike; you spot vanishing gradients when the loss plateaus early and the first layers' weights hardly move. The cures in this lesson — init, normalisation, clipping, and warmup — all work by keeping that running product near 1.

2 Reading the Loss Curve: spikes, NaN, and Inf

The loss numbers tell you almost everything. A sudden jump (say 3x the previous step) is a loss spike — often a too-high learning rate or an unclipped gradient. A NaN ("not a number") means a value became undefined and the run is dead from that point on. The neat trick: NaN is the only value that is not equal to itself, so loss != loss is a one-line NaN test (no import needed).

Run this and watch it flag the spike at step 4 and the NaN at step 7:

# Worked example: read a loss curve and spot the trouble
# A training run that goes wrong usually leaves clues in the loss numbers.

losses = [2.40, 2.10, 1.85, 1.60, 9.80, 1.55, 1.50, float("nan"), 1.48]
#                                  ^^^^ a loss SPIKE         ^^^ then NaN

prev = losses[0]
for step, loss in enumerate(losses):
    note = ""
    if loss != loss:                 # NaN is the only value not equal to itself
        note = "<-- NaN! training is broken from here"
    elif loss > prev * 3:            # 3x jump = a loss spike
        note = "<-- loss SPIKE (3x the previous step)"
    print(f"step {step}: loss={loss}  {note}")
    if loss == loss:                 # only update prev on real numbers
        prev = loss

# Expected output:
# step 0: loss=2.4
# step 1: loss=2.1
# step 2: loss=1.85
# step 3: loss=1.6
# step 4: loss=9.8  <-- loss SPIKE (3x the previous step)
# step 5: loss=1.55
# step 6: loss=1.5
# step 7: loss=nan  <-- NaN! training is broken from here
# step 8: loss=1.48

3 Gradient Clipping by Norm

The cleanest fix for exploding gradients is to cap the norm (the length) of the gradient. If the gradient's L2 norm exceeds a threshold, you scale the whole vector down by one factor — max_norm / norm — so its direction stays the same and only its size is capped. That last point is why clipping by norm beats clipping each value independently: distorting direction sends the optimiser somewhere it never intended to go.

This worked example clips the vector [3, 4, 12] (norm 13) down to norm 5:

# Worked example: clip a gradient vector to a maximum norm
# Exploding gradients = the update vector gets huge and blows up the weights.
# The fix: if its length (L2 norm) is too big, scale the WHOLE vector down.

def l2_norm(vec):
    """Length of the vector = sqrt(sum of squares)."""
    return sum(v * v for v in vec) ** 0.5

def clip_grad_norm(grad, max_norm):
    """Scale grad so its norm is at most max_norm. Direction is preserved."""
    norm = l2_norm(grad)
    if norm > max_norm:
        scale = max_norm / norm          # shrink factor, same for every element
        return [g * scale for g in grad], norm, scale
    return list(grad), norm, 1.0         # already small enough: leave it alone

grad = [3.0, 4.0, 12.0]                  # norm = sqrt(9+16+144) = 13.0
clipped, before, scale = clip_grad_norm(grad, max_norm=5.0)

print(f"norm before : {before:.2f}")
print(f"scale factor : {scale:.4f}")
print(f"clipped      : {[round(c, 3) for c in clipped]}")
print(f"norm after   : {l2_norm(clipped):.2f}")

# Expected output:
# norm before : 13.00
# scale factor : 0.3846
# clipped      : [1.154, 1.538, 4.615]
# norm after   : 5.00

In real code you never hand-roll this — PyTorch's clip_grad_norm_ does it across every parameter at once and returns the pre-clip norm so you can log it:

import torch
import torch.nn as nn

# The real thing: PyTorch ships clip_grad_norm_ so you never hand-roll it.
# It clips across ALL the model's parameters at once, in place.

torch.manual_seed(0)                     # reproducibility: same run every time
model = nn.Linear(4, 1)
x = torch.randn(8, 4)
y = torch.randn(8, 1)

loss = ((model(x) - y) ** 2).mean()
loss.backward()                          # fills .grad on every parameter

before = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
# clip_grad_norm_ RETURNS the total norm BEFORE clipping, then clips in place.
after = sum(p.grad.norm() ** 2 for p in model.parameters()) ** 0.5

print(f"total grad norm before clip: {before:.4f}")
print(f"total grad norm after  clip: {after:.4f}")   # <= 1.0

# Expected output:
# total grad norm before clip: 1.8262
# total grad norm after  clip: 1.0000

🎯 Your Turn: finish the norm clipper

Fill in the two blanks so the clipper scales [6, 8] (norm 10) down to norm 5.

# 🎯 YOUR TURN — finish the gradient-norm clipper (fill in the ___)

def l2_norm(vec):
    return sum(v * v for v in vec) ** 0.5

def clip_grad_norm(grad, max_norm):
    norm = l2_norm(grad)
    if norm > max_norm:
        scale = ___                       # 👉 shrink factor = max_norm / norm
        return [g * scale for g in grad]
    return list(grad)                     # already small enough

grad = [6.0, 8.0]                         # norm = sqrt(36+64) = 10.0
clipped = clip_grad_norm(grad, max_norm=___)   # 👉 clip to a max norm of 5.0
print("clipped :", [round(c, 3) for c in clipped])
print("new norm:", round(l2_norm(clipped), 2))

# ✅ Expected output:
# clipped : [3.0, 4.0]
# new norm: 5.0

4 Learning-Rate Warmup, Schedules, and Mixed Precision

A fresh model has random weights, so the first gradients are loud. Warmup ramps the learning rate linearly from near zero up to your target over the first few steps, letting Adam's running statistics settle before you take big steps. After warmup you usually decay the rate — cosine or step schedules — for a cleaner final model.

Mixed precision (AMP) speeds training by computing in lower precision, but tiny gradients can underflow to zero. GradScaler multiplies the loss up before backward() and unscales afterwards; it also skips the step automatically when it detects inf/NaN gradients. Always unscale_ before clipping so the clip sees the true norms:

import torch
import torch.nn as nn
from torch.cuda.amp import autocast, GradScaler

# Production stability stack: warmup schedule + AMP loss scaling + clipping.
torch.manual_seed(0)

model = nn.Linear(4, 1)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
scaler = GradScaler()                    # AMP: scales loss so small grads don't underflow to 0
warmup_steps = 5
base_lr = 1e-3

def lr_at(step):
    # Linear warmup: ramp 0 -> base_lr over warmup_steps, then hold.
    if step < warmup_steps:
        return base_lr * (step + 1) / warmup_steps
    return base_lr

for step in range(8):
    for group in opt.param_groups:
        group["lr"] = lr_at(step)        # apply the schedule before each step
    opt.zero_grad()
    with autocast():                     # run forward in lower precision (fast)
        loss = ((model(torch.randn(8, 4)) - 1.0) ** 2).mean()
    scaler.scale(loss).backward()        # scale UP before backward
    scaler.unscale_(opt)                 # unscale BEFORE clipping (clip sees true norms)
    nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    scaler.step(opt)                     # skips the step automatically if grads are inf/nan
    scaler.update()
    print(f"step {step}: lr={lr_at(step):.5f}")

# Expected output:
# step 0: lr=0.00020
# step 1: lr=0.00040
# step 2: lr=0.00060
# step 3: lr=0.00080
# step 4: lr=0.00100
# step 5: lr=0.00100
# step 6: lr=0.00100
# step 7: lr=0.00100

5 Weight Init and Normalisation for Stability

Initialisation sets the starting scale of every weight so the signal neither shrinks nor blows up as it passes through layers. Use He (Kaiming) init for ReLU networks and Xavier (Glorot) for tanh/sigmoid. Never use plain randn * 1.0 or randn * 0.01 for a deep network — the first explodes, the second vanishes.

Normalisation re-centres each layer's activations during training so the next layer always sees a well-behaved distribution. Batch Norm normalises across the batch (great for CNNs with large batches); Layer Norm normalises across features (the default for Transformers and RNNs, and the safe choice for small batches).

Use for ReLU / large batches

nn.init.kaiming_normal_(w)   # He
nn.BatchNorm2d(channels)     # Batch Norm

Use for Transformers / small batches

nn.init.xavier_uniform_(w)   # Xavier
nn.LayerNorm(features)       # Layer Norm

🎯 Your Turn: find the first bad step

Detecting trouble early is half the battle. Fill in the blanks so the loop reports the first NaN or Inf in the loss list.

# 🎯 YOUR TURN — report the FIRST broken step in a loss list (fill in the ___)
import math

losses = [1.9, 1.5, 1.2, float("inf"), 0.9, float("nan")]

first_bad = None
for step, loss in enumerate(losses):
    if math.isnan(loss) or ___:           # 👉 also catch infinity: math.isinf(loss)
        first_bad = step
        break                             # stop at the first bad step

if first_bad is ___:                      # 👉 None means we found nothing bad
    print("All losses are finite — training looks healthy")
else:
    print(f"First bad loss at step {first_bad}: {losses[first_bad]}")

# ✅ Expected output:
# First bad loss at step 3: inf

6 Reproducibility: seed everything

If two runs give different results you can't tell whether a change helped or you just got lucky. Seed every source of randomness before you build the model — Python's random, NumPy, and PyTorch (CPU and GPU):

import random, numpy as np, torch

def seed_everything(seed=42):
    random.seed(seed)                 # Python's RNG
    np.random.seed(seed)              # NumPy
    torch.manual_seed(seed)           # PyTorch CPU
    torch.cuda.manual_seed_all(seed)  # PyTorch GPU
    # For bitwise-identical runs (slower):
    torch.use_deterministic_algorithms(True)

seed_everything(42)   # call ONCE, at the very top, before building the model

Call seed_everything once at the top of your script. Without it, weight init, shuffling, and dropout all differ between runs.

🎯 Mini-Challenge: a tiny training guard

Support is faded now — only a comment outline is given. Write the whole function yourself.

# 🎯 MINI-CHALLENGE: a tiny "training guard"
# Write a function guard(grad, loss, max_norm) that:
#   1. Returns "skip" if loss is NaN or Inf   (use math.isnan / math.isinf)
#   2. Otherwise clips grad to max_norm by L2 norm and returns the clipped list
# Then test it on a healthy step AND a NaN step.
#
# ✅ Expected output:
# healthy: [3.0, 4.0]
# broken : skip

import math

# your code here

Common Errors (And How to Fix Them)

❌ No gradient clipping → loss explodes to NaN

An unclipped gradient spikes, the weights jump, and the loss goes inf then NaN.

✅ Fix: add nn.utils.clip_grad_norm_(model.parameters(), 1.0) after backward() and before optimizer.step().

❌ Learning rate too high → loss spikes or diverges

The loss oscillates or climbs instead of falling — the steps are too big.

✅ Fix: lower the LR (try 10x smaller) and add warmup so early steps stay small.

❌ Bad initialisation → dead neurons or instant explosion

randn * 1.0 explodes through deep layers; randn * 0.01 makes activations vanish and ReLUs go dead.

✅ Fix: use He init for ReLU (kaiming_normal_), Xavier for tanh/sigmoid.

❌ Not seeding → results change every run

You can't reproduce a result or tell whether a tweak actually helped.

✅ Fix: call seed_everything(42) once at the very top, before building the model.

❌ Ignoring NaN → every later step is garbage

Once a NaN enters the weights it poisons all subsequent updates, but training keeps "running".

✅ Fix: check each step with loss != loss or math.isnan(loss) and stop or skip immediately.

📋 Quick Reference

TechniqueFixesHow / Use With
Clip by normExploding gradientsclip_grad_norm_(params, 1.0)
LR warmupEarly instabilityRamp 0 → base over first steps
LR scheduleNoisy late trainingCosine / step decay after warmup
He initVanishing activationsReLU networks
Xavier initVanishing activationstanh, sigmoid
Batch NormCovariate shiftCNNs, large batches
Layer NormCovariate shiftTransformers, small batches
AMP + GradScalerUnderflow / inf gradsMixed-precision training
NaN checkSilent corruptionloss != loss / math.isnan
SeedingNon-reproducibilityseed_everything(42) once

🎉 Lesson Complete!

You can now keep a tall network from toppling: you read a loss curve for spikes and NaN, clip gradients to a safe norm (by hand and with clip_grad_norm_), warm up and schedule the learning rate, pick init/normalisation that hold signals steady, and seed every RNG for reproducible runs.

Practice quiz

What causes exploding gradients during training?

  • Too few training examples
  • Using a learning rate that is too small
  • Gradients multiplied through many layers with factors above 1 on average grow without bound
  • Normalising the inputs

Answer: Gradients multiplied through many layers with factors above 1 on average grow without bound. When per-layer factors exceed 1 on average, the product of gradients grows huge across layers — weights jump and the loss blows up to inf then NaN.

What characterises vanishing gradients?

  • Gradients shrink toward zero across layers, so early layers barely update
  • The loss becomes negative
  • The model trains too fast
  • The batch size is too large

Answer: Gradients shrink toward zero across layers, so early layers barely update. With factors below 1, the gradient product shrinks toward zero, so the earliest layers receive almost no signal and stop learning.

Why is clipping gradients by norm preferred over clipping by value?

  • It is faster to compute
  • It removes the need for a learning rate
  • It only affects the largest element
  • It scales the whole vector by one factor, preserving direction and only capping magnitude

Answer: It scales the whole vector by one factor, preserving direction and only capping magnitude. Clipping by value distorts the gradient's direction; clipping by global norm scales the whole vector once, preserving direction and capping only its length.

When a gradient's norm exceeds max_norm, what scale factor is applied?

  • norm / max_norm
  • max_norm / norm
  • max_norm * norm
  • 1 / max_norm

Answer: max_norm / norm. The scale factor is max_norm / norm, applied to every element, so the clipped vector has length exactly max_norm with its direction unchanged.

Why is learning-rate warmup useful?

  • It ramps the LR up gradually so early noisy gradients (and Adam's stats) settle before big steps
  • It permanently increases the learning rate
  • It removes the need for gradient clipping
  • It shuffles the training data

Answer: It ramps the LR up gradually so early noisy gradients (and Adam's stats) settle before big steps. Fresh weights give large noisy gradients; ramping the LR from near zero lets Adam's running statistics settle, preventing an early loss explosion.

What is a one-line way to detect a NaN loss with no imports?

  • loss > 0
  • loss == 0
  • loss != loss
  • loss < float('inf')

Answer: loss != loss. NaN is the only value not equal to itself, so loss != loss is True exactly when loss is NaN — a no-import NaN check.

Which weight initialisation is recommended for ReLU networks?

  • Xavier (Glorot)
  • He (Kaiming)
  • All zeros
  • randn * 1.0

Answer: He (Kaiming). He (Kaiming) init suits ReLU networks; Xavier (Glorot) suits tanh/sigmoid. Plain randn*1.0 explodes and randn*0.01 vanishes in deep nets.

Which normalisation is the default choice for Transformers and small batches?

  • Batch Norm
  • No normalisation
  • Dropout
  • Layer Norm

Answer: Layer Norm. Layer Norm normalises across features and is the safe default for Transformers, RNNs, and small batches; Batch Norm suits CNNs with large batches.

In mixed-precision training, what does GradScaler do when it detects inf/NaN gradients?

  • It crashes the program
  • It skips the optimizer step automatically for that iteration
  • It increases the learning rate
  • It converts the model to full precision permanently

Answer: It skips the optimizer step automatically for that iteration. GradScaler scales the loss up before backward to prevent underflow, and automatically skips the step when it sees inf or NaN gradients.

Why should you seed every RNG source before training?

  • To make the model train faster
  • To reduce memory usage
  • So runs are reproducible and you can tell whether a change actually helped
  • To avoid gradient clipping

Answer: So runs are reproducible and you can tell whether a change actually helped. Seeding random, NumPy, and PyTorch (CPU and GPU) makes runs reproducible, so improvements reflect real changes rather than random noise between runs.

Continue this course

Frequently asked questions

What causes exploding and vanishing gradients?

Both come from repeatedly multiplying gradients through many layers during backpropagation. If the factors are bigger than 1 on average, the product grows without bound (exploding — you see huge values, then NaN). If they are smaller than 1, the product shrinks toward zero (vanishing — early layers stop learning). Good weight initialisation, normalisation, residual connections, and gradient clipping all keep these products near 1.

Should I clip gradients by value or by norm?

Clip by norm. Clipping by value caps each element independently, which distorts the gradient's direction. Clipping by global norm scales the whole vector by one factor only when its total length exceeds the threshold, so the direction is preserved and only the magnitude is capped. In PyTorch use torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm). Typical max_norm is 1.0 for Transformers and around 5.0 for RNNs.

Why do I need learning-rate warmup?

At the very start of training the weights are random, so the first few gradients can be large and noisy — especially with adaptive optimisers like Adam, whose running statistics are not yet warmed up. Ramping the learning rate linearly from near zero to the target over the first few hundred to a few thousand steps lets those statistics settle before you take big steps, which prevents an early loss explosion. After warmup most people decay the rate (cosine or step schedule) for a cleaner final model.

My loss became NaN — what now?

NaN means a number became undefined (often 0/0, log(0), or overflow from an exploding gradient). Do not ignore it: once a NaN enters the weights every later step is garbage. Detect it early (a loss is NaN when loss != loss, or use math.isnan), then fix the cause — lower the learning rate, add or tighten gradient clipping, check for log(0) or divide-by-zero in your loss, and confirm your inputs contain no NaN/Inf. With mixed precision, GradScaler skips the update automatically when it sees inf or NaN gradients.

How do I make a training run reproducible?

Seed every source of randomness before you build the model: random.seed(s), numpy.random.seed(s), and torch.manual_seed(s) (plus torch.cuda.manual_seed_all(s) on GPU). For bitwise-identical runs also set torch.use_deterministic_algorithms(True) and avoid nondeterministic GPU kernels. Reproducibility is what lets you tell whether a change actually helped, instead of chasing random noise between runs.

Related lessons