Building Models with nn.Module
Weight initialisation
The random numbers a network starts with decide whether the signal survives fifty layers or dies on the way — and PyTorch's defaults are already chosen with that in mind.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Weight initialisation is choosing the random starting values for a network's numbers, so that learning can get going at all.
Imagine a message passed along a line of fifty people. If everyone whispers a little softer than they heard, the message is silence by person twenty. If everyone shouts a little louder, it is a distorted roar. The message survives only if each person repeats it at roughly the volume they heard.
A deep network is that line of people. Each layer passes the signal on, a bit louder or softer depending on its starting numbers. Initialisation is picking a starting volume that keeps the message alive all the way down.
Why it exists
Training starts from randomness — the network has learned nothing yet. But "random" has a size, and the size is everything.
Start too small, and the signal shrinks layer by layer until nothing reaches the end. The training signal coming back is equally dead, so nothing learns. Start too big, and each layer amplifies until the numbers saturate or blow up.
For years this was a real wall. Deep networks would not train, and starting volume was one of the main reasons. The fix was worked out around 2010: match the starting size to how many inputs each layer receives.
How it works
too quiet: signal → layer → layer → layer → ... → nothing
too loud: signal → layer → layer → layer → ... → distorted roar
matched: signal → layer → layer → layer → ... → signal, still aliveThe matched size depends on the layer's width and on which activation function follows it. PyTorch layers arrive with a sensible matched size already applied.
A real example you have seen
Every large model you have used — translation, photo search, chatbots — began as pure noise. Its abilities exist because that noise was sized correctly enough for training to take hold.
Remember this
- Networks start random, and the size of that randomness decides whether training works.
- Too small starves the signal; too big saturates it.
- PyTorch's defaults are good. Change them when you know why, not before.
What to learn next
- model.train() and model.eval() — the mode switch every freshly initialised model must respect.
- Activation functions — why the starting size depends on the bend between layers.
- BatchNorm and its running statistics — the layer that made networks far less sensitive to initialisation.
Developer — Code and libraries.
Setup
pip install torchWritten and tested against torch 2.5 on CPU. Values below come from a fixed seed on that version — a different torch build may print slightly different numbers, with the same pattern.
Watch a signal live or die
import torch
torch.manual_seed(0)
def deep_stack(scale):
"""50 tanh layers with weights drawn at the given size."""
x = torch.randn(256, 512)
for _ in range(50):
w = torch.randn(512, 512) * scale
x = torch.tanh(x @ w)
saturated = (x.abs() > 0.99).float().mean().item()
return x.std().item(), saturated
for label, scale in [("too small (0.01)", 0.01),
("too big (0.30)", 0.30),
("xavier (1/sqrt(512))", 1 / 512**0.5)]:
std, sat = deep_stack(scale)
print(f"{label:24s} signal std {std:.4f} stuck at +/-1: {sat:.0%}")too small (0.01) signal std 0.0000 stuck at +/-1: 0% too big (0.30) signal std 0.9364 stuck at +/-1: 68% xavier (1/sqrt(512)) signal std 0.1006 stuck at +/-1: 0%
Read the middle line twice. The standard deviation looks healthy at 0.94 — but 68% of all values are jammed against the tanh ceiling. A saturated tanh passes almost no gradient, so that network is as untrainable as the silent one. Summary statistics can lie; check saturation too.
The Xavier size — one over the square root of the number of inputs — keeps the signal alive through all fifty layers. That is the whole trick.
Applying an initialisation on purpose
PyTorch's nn.Linear and nn.Conv2d already ship with a reasonable default. When you do want control, the pattern is apply:
import torch
from torch import nn
torch.manual_seed(0)
model = nn.Sequential(nn.Linear(64, 64), nn.ReLU(), nn.Linear(64, 10))
def init_weights(m):
if isinstance(m, nn.Linear):
nn.init.kaiming_normal_(m.weight, nonlinearity="relu")
nn.init.zeros_(m.bias)
model.apply(init_weights) # walks every submodule, calls the function
w = model[0].weight
print(f"first layer std: {w.std().item():.4f} (kaiming target: {(2/64)**0.5:.4f})")
print("bias all zero:", bool((model[0].bias == 0).all()))first layer std: 0.1808 (kaiming target: 0.1768) bias all zero: True
Kaiming initialisation is Xavier's sibling, tuned for ReLU: ReLU deletes half the signal (everything negative), so the weights start twice as large in variance to compensate. Use Kaiming with ReLU-family activations, Xavier with tanh or sigmoid.
model.apply(fn) visits every submodule recursively — the same register walk from the registration lesson. The isinstance check matters because apply also visits ReLU and the container itself.
Common mistakes
Initialising to zero. Every neuron in a layer then computes the identical thing and receives the identical correction, forever. The layer collapses to one effective neuron. Symmetry must be broken by randomness.
Initialising to a constant, or all the same sign. Same symmetry problem in disguise.
Re-initialising a pretrained model. apply(init_weights) after loading a checkpoint erases the training. Initialise first, load second — order matters.
Forgetting the seed when comparing runs. Two runs with different random starts differ for that reason alone. Set torch.manual_seed before creating the model when you compare anything.
Try it yourself
In signal_survival.py, swap torch.tanh for torch.relu and rerun. Watch the Xavier row die too — then fix it by scaling with sqrt(2/512) instead, which is exactly Kaiming's correction.
What to learn next
- model.train() and model.eval() — the mode switch every freshly initialised model must respect.
- Activation functions — why the starting size depends on the bend between layers.
- BatchNorm and its running statistics — the layer that made networks far less sensitive to initialisation.
Researcher — Mathematics and papers.
The variance argument
For a linear map $y = Wx$ with $W \in \mathbb{R}^{n_{out} \times n_{in}}$, entries i.i.d. with mean 0 and variance $\sigma^2$, and inputs i.i.d. with variance $v$:
$$ \mathrm{Var}(y_i) = n_{in} \, \sigma^2 \, v $$
Where $n_{in}$ is the fan-in (inputs per neuron), $\sigma^2$ the weight variance, and $v$ the input variance. Variance is preserved across the layer iff $\sigma^2 = 1/n_{in}$. Requiring the same for the backward pass gives $\sigma^2 = 1/n_{out}$; Glorot and Bengio (2010), Understanding the difficulty of training deep feedforward neural networks, split the difference:
$$ \sigma^2_{\text{Xavier}} = \frac{2}{n_{in} + n_{out}} $$
He et al. (2015), Delving Deep into Rectifiers, redid the computation for ReLU. Since ReLU zeroes the negative half, $\mathbb{E}[\mathrm{ReLU}(z)^2] = \tfrac{1}{2}\mathbb{E}[z^2]$ for symmetric $z$, so preserving variance needs a factor-2 correction:
$$ \sigma^2_{\text{Kaiming}} = \frac{2}{n_{in}} $$
Both derivations assume independence and zero mean at every layer — approximations that hold at initialisation and immediately break during training. They are statements about step zero only, which is exactly why they matter: step zero decides whether gradients exist at all.
PyTorch's actual default, a historical footnote
nn.Linear initialises with kaiming_uniform_(weight, a=math.sqrt(5)), which works out to $U(-1/\sqrt{n_{in}}, 1/\sqrt{n_{in}})$ — closer to the old LeCun scheme than to He's recommendation, kept for backward compatibility with the original Torch. It works well in practice; it is not the He (2015) formula, despite the function name. Bias is initialised uniformly from the same bound, not to zero. Verified against torch 2.5 source (torch/nn/modules/linear.py).
Beyond variance matching
- Orthogonal initialisation (Saxe et al., 2014, Exact solutions to the nonlinear dynamics of learning in deep linear networks): draw $W$ orthogonal so singular values are exactly 1; exact dynamical isometry for linear nets.
- LSUV (Mishkin and Matas, 2016): orthogonal draw, then rescale each layer empirically until unit output variance on real data.
- Fixup / zero-init residual branches (Zhang et al., 2019): residual nets train without normalisation if the last layer of each residual branch starts at zero — initialisation substituting for BatchNorm.
- muP (Yang et al., 2021, Tensor Programs V): parameterise init and learning rates so optimal hyperparameters transfer across width — initialisation as the lever for zero-shot hyperparameter transfer at scale.
The through-line: initialisation, normalisation, and architecture are partially interchangeable controls over the same quantity — signal propagation statistics at depth.
Reading
- Glorot and Bengio (2010); He et al. (2015) — the two formulas everyone uses.
- Saxe et al. (2014); Zhang et al. (2019); Yang et al. (2021) — what replaced guesswork at scale.
What to learn next
- model.train() and model.eval() — the mode switch every freshly initialised model must respect.
- Activation functions — why the starting size depends on the bend between layers.
- BatchNorm and its running statistics — the layer that made networks far less sensitive to initialisation.