Deep Learning with TensorFlow and PyTorch

Weight Initialization and Gradient Issues


You have to give the weights some starting value before training can begin. Zero seems tidy and unbiased. Try it.

Python
import torch.nn as nnmodel = nn.Sequential(nn.Linear(10, 64), nn.ReLU(), nn.Linear(64, 1))for p in model.parameters():    nn.init.zeros_(p)# Train for 500 epochs... the loss goes down slightly, then stops dead.

Here is why. Every unit in the hidden layer has identical weights, so every unit computes an identical output for any input. Every unit therefore receives an identical gradient, so every unit takes an identical update, so after the step they are still identical. The 64 units are one unit wearing 64 hats, forever. This is the symmetry problem, and no amount of training breaks it — the update rule preserves the symmetry exactly.

Fine: use random values. But which random values? This turns out to be a question with a precise answer, and getting it wrong produces a network that either explodes or falls silent before the first gradient ever arrives.

Gradient norm per layer, badly scaled init0.410.0860.0190.00410.0008801234layer 5 (last)layer 1(first)Each hop multiplies by roughly 0.2, so five layers back the signal is 500 times smaller.
The early layers are not learning slowly — they are receiving a gradient that arithmetic has already erased.

The experiment that shows what is at stake

Build a 10-layer network of 512 units each, feed it standardised input, and print the standard deviation of the activations at each layer. Nothing is trained; this is one forward pass.

Python
import torchdef probe(std, depth=10, width=512):    x = torch.randn(256, width)    out = []    for _ in range(depth):        W = torch.randn(width, width) * std        x = torch.relu(x @ W)        out.append(x.std().item())    return outfor name, s in [("std=1.0", 1.0), ("std=0.01", 0.01),                ("He: sqrt(2/512)", (2/512) ** 0.5)]:    stds = probe(s)    print(name, [f"{v:.2e}" for v in stds[::3]])
InitialisationLayer 1Layer 4Layer 7Layer 10Outcome
σ=1.0\sigma = 1.0~13∼5×104\sim 5 \times 10^{4}∼2×108\sim 2 \times 10^{8}∼8×1011\sim 8 \times 10^{11}Overflow, loss is nan
σ=0.01\sigma = 0.01~0.13∼6×10−4\sim 6 \times 10^{-4}∼3×10−6\sim 3 \times 10^{-6}∼1×10−8\sim 1 \times 10^{-8}Output is zero; no gradient
σ=2/512\sigma = \sqrt{2/512}~0.8~0.8~0.8~0.8Stable

The arithmetic behind those columns is simple. For z=∑i=1nwixiz = \sum_{i=1}^{n} w_i x_i with independent, zero-mean ww and xx:

Var(z)=n⋅Var(w)⋅Var(x)\text{Var}(z) = n \cdot \text{Var}(w) \cdot \text{Var}(x)

With n=512n = 512 and Var(w)=1\text{Var}(w) = 1, the variance is multiplied by 512 at every layer. Ten layers gives 51210≈1.2×1027512^{10} \approx 1.2 \times 10^{27}. With Var(w)=0.0001\text{Var}(w) = 0.0001, the multiplier is 512×0.0001=0.0512512 \times 0.0001 = 0.0512, and ten layers gives 10−1310^{-13}.

The design goal for initialisation is one sentence: choose the weight variance so that the activation variance is roughly preserved from layer to layer. Then nothing explodes and nothing vanishes, in either direction.

Xavier initialisation

Setting n⋅Var(w)=1n \cdot \text{Var}(w) = 1 gives Var(w)=1/nin\text{Var}(w) = 1/n_{\text{in}}, which preserves the variance of the forward signal. But gradients travel backwards through the transposed weight matrix, and by the same argument that direction wants Var(w)=1/nout\text{Var}(w) = 1/n_{\text{out}}. Both cannot hold unless the layer is square, so Xavier (also called Glorot) takes the harmonic compromise:

Var(w)=2nin+nout\text{Var}(w) = \frac{2}{n_{\text{in}} + n_{\text{out}}}

In the uniform form, weights are drawn from [−6/(nin+nout), +6/(nin+nout)]\left[-\sqrt{6/(n_{\text{in}}+n_{\text{out}})},\ +\sqrt{6/(n_{\text{in}}+n_{\text{out}})}\right] — the 6\sqrt{6} is there because a uniform distribution on [−a,a][-a, a] has variance a2/3a^2/3.

The derivation assumes the activation function is roughly linear near zero and does not change the variance. That is a fair description of tanh, and a poor one of ReLU.

He initialisation, and where the factor of 2 comes from

ReLU sets every negative value to zero. If the pre-activations are symmetric about zero, half of them are zeroed, and the variance of what survives is roughly half the variance of what went in:

Var(ReLU(z))≈12Var(z)\text{Var}(\text{ReLU}(z)) \approx \tfrac{1}{2}\text{Var}(z)

So a network using Xavier with ReLU loses half its variance at every layer. Ten layers gives a factor of 2−10≈0.0012^{-10} \approx 0.001 — a slow-motion version of the vanishing problem. Kaiming He's fix is to double the weight variance to compensate:

Var(w)=2nin\text{Var}(w) = \frac{2}{n_{\text{in}}}

Check it: Var(zout)=12⋅n⋅2n⋅Var(zin)=Var(zin)\text{Var}(z_{\text{out}}) = \frac{1}{2} \cdot n \cdot \frac{2}{n} \cdot \text{Var}(z_{\text{in}}) = \text{Var}(z_{\text{in}}). Preserved exactly, which is the third row of the table above.

ActivationUseVariancePyTorch
ReLU, Leaky ReLU, GELUHe / Kaiming2/nin2/n_{\text{in}}nn.init.kaiming_normal_(w, nonlinearity='relu')
tanh, sigmoidXavier / Glorot2/(nin+nout)2/(n_{\text{in}}+n_{\text{out}})nn.init.xavier_normal_(w)
SELULeCun1/nin1/n_{\text{in}}nn.init.normal_(w, std=(1/fan_in)**0.5)
Any — biasesZeros—nn.init.zeros_(b)

Biases can safely start at zero: the symmetry problem is already broken by the random weights, so there is nothing to break.

What the frameworks do by default

Keras Dense defaults to Glorot uniform. PyTorch nn.Linear defaults to a Kaiming-uniform variant with a=5a=\sqrt{5}, which works out close to 1/(3nin)1/(3 n_{\text{in}}) — noticeably smaller than true He initialisation. Neither default is wrong exactly, but neither is ideal for a deep ReLU stack, and setting it explicitly costs four lines:

Python
import torch.nn as nndef init_weights(m):    if isinstance(m, (nn.Linear, nn.Conv2d)):        nn.init.kaiming_normal_(m.weight, mode="fan_in", nonlinearity="relu")        if m.bias is not None:            nn.init.zeros_(m.bias)model.apply(init_weights)      # applies recursively to every submodule
Python
from tensorflow.keras import layerslayers.Dense(256, activation="relu",             kernel_initializer="he_normal",             bias_initializer="zeros")

Two situations where initialisation genuinely decides whether the model trains at all: very deep networks without normalisation layers, and recurrent networks, where an orthogonal initialisation of the recurrent weight matrix (nn.init.orthogonal_) keeps repeated multiplication by the same matrix from exploding or collapsing over long sequences.

Vanishing gradients

The backward pass multiplies the error signal at every layer by the weights and by the activation's derivative. Over LL layers you get a product of LL such factors, and products of numbers below one shrink geometrically.

With sigmoid activations, the derivative σ′(z)=σ(z)(1−σ(z))\sigma'(z) = \sigma(z)(1-\sigma(z)) peaks at 0.25 and is usually far lower:

DepthBest-case factor 0.25L0.25^LGradient reaching layer 1
59.8×10−49.8 \times 10^{-4}A thousandth
109.5×10−79.5 \times 10^{-7}A millionth
209.1×10−139.1 \times 10^{-13}Below float32 precision in practice

What it looks like in a real run: the loss falls for a few epochs, then flattens well above where it should. Later layers' weights change; early layers' weights are nearly identical to their initial values. The model is effectively shallow — only the last few layers are learning.

RemedyMechanism
ReLU-family activationsDerivative is exactly 1 for positive inputs, so no shrinking factor
He initialisationKeeps the weight-matrix factor near 1 rather than below it
Batch or layer normalisationRe-standardises activations at every layer, bounding how far they can drift
Residual connectionsGives the gradient an unobstructed path that skips the multiplications entirely

Exploding gradients

The same product, with factors above one. 1.520≈33001.5^{20} \approx 3300; 220≈1062^{20} \approx 10^6. The weight update overshoots so far that the loss jumps to an enormous value, and one more step produces inf and then nan. Recurrent networks are especially prone, because the same matrix is applied at every timestep, so a spectral radius slightly above 1 compounds over hundreds of steps.

The standard remedy is blunt and effective: rescale the whole gradient whenever its norm exceeds a threshold.

Python
for xb, yb in train_loader:    opt.zero_grad()    loss = loss_fn(model(xb), yb)    loss.backward()    # Rescale ALL gradients together so the global norm is at most 1.0.    # Direction is preserved; only the magnitude is capped.    total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)    opt.step()

Clip by norm, not by value. clip_grad_value_ caps each component independently, which changes the direction of the gradient vector — you are no longer descending the loss, you are descending something adjacent to it. Norm clipping scales the whole vector, so the direction is untouched.

clip_grad_norm_ returns the norm before clipping, which is worth logging. If it sits comfortably below your threshold, clipping is doing nothing and you can stop worrying. If it is being clipped on most steps, your learning rate is probably too high and clipping is masking the symptom.

Python
opt = tf.keras.optimizers.Adam(learning_rate=1e-3, global_clipnorm=1.0)

Residual connections, and why they changed everything

A residual block computes y=x+F(x)y = x + F(x) instead of y=F(x)y = F(x). Differentiate it:

∂y∂x=1+∂F∂x\frac{\partial y}{\partial x} = 1 + \frac{\partial F}{\partial x}

That standalone 11 is the entire point. Even if ∂F/∂x\partial F/\partial x is tiny, the derivative through the block is at least 1, so the gradient passes through undiminished. Stack fifty such blocks and the gradient still arrives at the first layer with useful magnitude — which is why residual networks reach hundreds of layers where plain stacks stall around twenty.

Python
class ResidualBlock(nn.Module):    def __init__(self, dim):        super().__init__()        self.fc1 = nn.Linear(dim, dim, bias=False)        self.bn1 = nn.BatchNorm1d(dim)        self.fc2 = nn.Linear(dim, dim, bias=False)        self.bn2 = nn.BatchNorm1d(dim)    def forward(self, x):        out = torch.relu(self.bn1(self.fc1(x)))        out = self.bn2(self.fc2(out))        return torch.relu(out + x)          # add the input, THEN activate

Two things to get right. The addition must come before the final activation — putting the ReLU inside the branch means the skip path passes through a non-linearity and the clean gradient route is lost. And if the block changes the dimension, the shortcut needs a projection (nn.Linear(dim_in, dim_out)), or the shapes will not add.

A useful trick from modern practice: initialise the final normalisation layer's γ\gamma to zero. Then F(x)=0F(x) = 0 at the start and the block is exactly the identity, so a very deep network begins as a shallow one and grows into its depth as training proceeds. It measurably stabilises the first few hundred steps.

Measuring gradient flow directly

All of the above becomes concrete the moment you print the per-layer gradient norms. This is the single most useful diagnostic in deep learning and it takes six lines.

Python
def gradient_report(model):    print(f"{'layer':30s} {'grad norm':>12s} {'weight norm':>12s} {'ratio':>10s}")    for name, p in model.named_parameters():        if p.grad is None or p.ndim < 2:            continue        g, w = p.grad.norm().item(), p.data.norm().item()        print(f"{name:30s} {g:12.3e} {w:12.3e} {g/max(w,1e-12):10.2e}")loss.backward()gradient_report(model)
What you seeDiagnosisAction
Norms roughly similar across all layersHealthyNothing
Norms shrink by 10× or more per layer going backwardsVanishingHe init, ReLU, normalisation, residual connections
Norms above 10310^{3} anywhereExplodingClip to 1.0; lower the learning rate
Exactly zero for a whole layerDisconnected, frozen, or all-dead ReLUsCheck requires_grad; check the forward path reaches this layer
None for a parameterNever entered the computation graphLook for a .detach(), a NumPy round-trip, or an unused layer

The ratio column is the most informative. Gradient norm divided by weight norm tells you the relative size of the update each layer is about to receive. A healthy value is somewhere around 10−310^{-3} to 10−210^{-2} — the weights change by roughly a tenth of a percent to a percent per step. Ratios of 10−710^{-7} mean that layer is frozen in all but name; ratios near 1 mean the layer is being completely rewritten every step and training will not converge.

A configuration that will not fight you

For anything you build today, the combination below removes initialisation and gradient flow from the list of things that can go wrong, leaving you free to worry about the model itself.

Python
model = nn.Sequential(    nn.Linear(n_in, 512, bias=False), nn.BatchNorm1d(512), nn.ReLU(),    ResidualBlock(512),    ResidualBlock(512),    nn.Linear(512, n_out),)model.apply(init_weights)                 # He normal, zero biasesopt = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-2)sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=1e-3,                                            total_steps=epochs * len(train_loader))for xb, yb in train_loader:    opt.zero_grad()    loss = loss_fn(model(xb), yb)    loss.backward()    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)    opt.step()    sched.step()

He initialisation because the activations are ReLU. Normalisation to keep activations in range. Residual connections so gradients reach the bottom. Clipping as a cheap insurance policy against a single bad batch. A warm-up-then-decay schedule so the first steps, taken from a random starting point, are small.

When something still goes wrong, run the gradient report before you change anything else. It converts a vague symptom — "the loss plateaued" — into a specific one — "layers 1 through 4 have gradient norms of 10−910^{-9}" — and specific symptoms have specific fixes. Guessing at hyperparameters when you could be reading the actual numbers is the most common way to lose a day.