Course Content
Deep Learning with TensorFlow and PyTorch
4 sections · 15 lessons
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.
1import torch.nn as nn23model = nn.Sequential(nn.Linear(10, 64), nn.ReLU(), nn.Linear(64, 1))4for p in model.parameters():5 nn.init.zeros_(p)6# 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.
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.
1import torch23def probe(std, depth=10, width=512):4 x = torch.randn(256, width)5 out = []6 for _ in range(depth):7 W = torch.randn(width, width) * std8 x = torch.relu(x @ W)9 out.append(x.std().item())10 return out1112for name, s in [("std=1.0", 1.0), ("std=0.01", 0.01),13 ("He: sqrt(2/512)", (2/512) ** 0.5)]:14 stds = probe(s)15 print(name, [f"{v:.2e}" for v in stds[::3]])| Initialisation | Layer 1 | Layer 4 | Layer 7 | Layer 10 | Outcome |
|---|---|---|---|---|---|
| σ=1.0 | ~13 | ∼5×104 | ∼2×108 | ∼8×1011 | Overflow, loss is nan |
| σ=0.01 | ~0.13 | ∼6×10−4 | ∼3×10−6 | ∼1×10−8 | Output is zero; no gradient |
| σ=2/512 | ~0.8 | ~0.8 | ~0.8 | ~0.8 | Stable |
The arithmetic behind those columns is simple. For z=∑i=1nwixi with independent, zero-mean w and x:
With n=512 and Var(w)=1, the variance is multiplied by 512 at every layer. Ten layers gives 51210≈1.2×1027. With Var(w)=0.0001, the multiplier is 512×0.0001=0.0512, and ten layers gives 10−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)=1 gives Var(w)=1/nin, 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. Both cannot hold unless the layer is square, so Xavier (also called Glorot) takes the harmonic compromise:
In the uniform form, weights are drawn from [−6/(nin+nout), +6/(nin+nout)] — the 6 is there because a uniform distribution on [−a,a] has variance a2/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:
So a network using Xavier with ReLU loses half its variance at every layer. Ten layers gives a factor of 2−10≈0.001 — a slow-motion version of the vanishing problem. Kaiming He's fix is to double the weight variance to compensate:
Check it: Var(zout)=21⋅n⋅n2⋅Var(zin)=Var(zin). Preserved exactly, which is the third row of the table above.
| Activation | Use | Variance | PyTorch |
|---|---|---|---|
| ReLU, Leaky ReLU, GELU | He / Kaiming | 2/nin | nn.init.kaiming_normal_(w, nonlinearity='relu') |
| tanh, sigmoid | Xavier / Glorot | 2/(nin+nout) | nn.init.xavier_normal_(w) |
| SELU | LeCun | 1/nin | nn.init.normal_(w, std=(1/fan_in)**0.5) |
| Any — biases | Zeros | — | 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=5, which works out close to 1/(3nin) — 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:
1import torch.nn as nn23def init_weights(m):4 if isinstance(m, (nn.Linear, nn.Conv2d)):5 nn.init.kaiming_normal_(m.weight, mode="fan_in", nonlinearity="relu")6 if m.bias is not None:7 nn.init.zeros_(m.bias)89model.apply(init_weights) # applies recursively to every submodule1from tensorflow.keras import layers23layers.Dense(256, activation="relu",4 kernel_initializer="he_normal",5 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 L layers you get a product of L such factors, and products of numbers below one shrink geometrically.
With sigmoid activations, the derivative σ′(z)=σ(z)(1−σ(z)) peaks at 0.25 and is usually far lower:
| Depth | Best-case factor 0.25L | Gradient reaching layer 1 |
|---|---|---|
| 5 | 9.8×10−4 | A thousandth |
| 10 | 9.5×10−7 | A millionth |
| 20 | 9.1×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.
| Remedy | Mechanism |
|---|---|
| ReLU-family activations | Derivative is exactly 1 for positive inputs, so no shrinking factor |
| He initialisation | Keeps the weight-matrix factor near 1 rather than below it |
| Batch or layer normalisation | Re-standardises activations at every layer, bounding how far they can drift |
| Residual connections | Gives the gradient an unobstructed path that skips the multiplications entirely |
Exploding gradients
The same product, with factors above one. 1.520≈3300; 220≈106. 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.
1for xb, yb in train_loader:2 opt.zero_grad()3 loss = loss_fn(model(xb), yb)4 loss.backward()56 # Rescale ALL gradients together so the global norm is at most 1.0.7 # Direction is preserved; only the magnitude is capped.8 total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)910 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.
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) instead of y=F(x). Differentiate it:
That standalone 1 is the entire point. Even if ∂F/∂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.
1class ResidualBlock(nn.Module):2 def __init__(self, dim):3 super().__init__()4 self.fc1 = nn.Linear(dim, dim, bias=False)5 self.bn1 = nn.BatchNorm1d(dim)6 self.fc2 = nn.Linear(dim, dim, bias=False)7 self.bn2 = nn.BatchNorm1d(dim)89 def forward(self, x):10 out = torch.relu(self.bn1(self.fc1(x)))11 out = self.bn2(self.fc2(out))12 return torch.relu(out + x) # add the input, THEN activateTwo 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 γ to zero. Then F(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.
1def gradient_report(model):2 print(f"{'layer':30s} {'grad norm':>12s} {'weight norm':>12s} {'ratio':>10s}")3 for name, p in model.named_parameters():4 if p.grad is None or p.ndim < 2:5 continue6 g, w = p.grad.norm().item(), p.data.norm().item()7 print(f"{name:30s} {g:12.3e} {w:12.3e} {g/max(w,1e-12):10.2e}")89loss.backward()10gradient_report(model)| What you see | Diagnosis | Action |
|---|---|---|
| Norms roughly similar across all layers | Healthy | Nothing |
| Norms shrink by 10× or more per layer going backwards | Vanishing | He init, ReLU, normalisation, residual connections |
| Norms above 103 anywhere | Exploding | Clip to 1.0; lower the learning rate |
| Exactly zero for a whole layer | Disconnected, frozen, or all-dead ReLUs | Check requires_grad; check the forward path reaches this layer |
None for a parameter | Never entered the computation graph | Look 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−3 to 10−2 — the weights change by roughly a tenth of a percent to a percent per step. Ratios of 10−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.
1model = nn.Sequential(2 nn.Linear(n_in, 512, bias=False), nn.BatchNorm1d(512), nn.ReLU(),3 ResidualBlock(512),4 ResidualBlock(512),5 nn.Linear(512, n_out),6)7model.apply(init_weights) # He normal, zero biases89opt = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-2)10sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=1e-3,11 total_steps=epochs * len(train_loader))1213for xb, yb in train_loader:14 opt.zero_grad()15 loss = loss_fn(model(xb), yb)16 loss.backward()17 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)18 opt.step()19 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−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.