Um momento
0x90Lesson 10 of 17

Initialization, normalization and residuals

Keep signals and gradients healthy through deep networks with good initialization, batch and layer norm, and skip connections.

30 min 7-question quiz 2 code exercises
By the end of this lesson you can
  • Explain vanishing and exploding activations and gradients in deep networks
  • Initialize weights with Xavier or He scaling
  • Describe how batch norm, layer norm and residual connections make deep networks trainable

Stack enough layers and something alarming happens. Each layer multiplies the signal by its weights; if those multiplications shrink it a little, after 50 layers it’s vanished; if they grow it a little, it explodes. The same happens to gradients on the way back. Either way, training stalls or breaks.

Watch it happen to a batch of random inputs passing through 50 ReLU layers of 256 units:

signal_through_depth.py
1import numpy as np
2
3def spread_after(depth, weight_std, width=256):
4    rng = np.random.default_rng(0)
5    x = rng.normal(size=(100, width))
6    for _ in range(depth):
7        W = rng.normal(0, weight_std, size=(width, width))
8        x = np.maximum(0, x @ W)
9    return x.std()
10
11he = np.sqrt(2 / 256)
12for name, std in [("too small (0.01)", 0.01), ("too big (0.1)", 0.1), (f"He ({he:.3f})", he)]:
13    print(f"{name:<16}", " ".join(f"{spread_after(depth, std):9.2e}" for depth in (1, 10, 50)))
Output
too small (0.01)  9.25e-02  3.32e-10  2.87e-48
too big (0.1)     9.25e-01  3.32e+00  2.87e+02
He (0.088)        8.18e-01  9.67e-01  6.00e-01

With tiny weights the signal is gone by layer 10. With weights only slightly too big, it grows layer after layer and explodes - 300 times larger by layer 50. He initialization keeps it roughly the same size at every depth. The rule of thumb is to scale the random weights by the number of inputs to the layer, ninn_{in}:

InitializationWeight standard deviationUse with
Xavier / Glorot (2010)1/nin\sqrt{1 / n_{in}}tanh, sigmoid
He / Kaiming (2015)2/nin\sqrt{2 / n_{in}}ReLU and friends - the extra 2 makes up for ReLU zeroing half the signal

Frameworks pick sensible defaults (PyTorch’s nn.Linear uses a Kaiming-style uniform init), but when you build a layer yourself, this is the line that decides whether deep training works at all. Biases usually start at zero.

Normalization layers

Good initialization helps at the start, but the weights change as training goes on. Normalization layers keep activations well-scaled the whole time: standardize to mean 0 and variance 1, then let the network re-scale with learned parameters γ\gamma (scale) and β\beta (shift):

x^=x−μσ2+ϵ,y=γx^+β\hat{x} = \frac{x - \mu}{\sqrt{\sigma^2 + \epsilon}}, \qquad y = \gamma \hat{x} + \beta

  • Batch normalization (2015) computes μ and σ for each feature across the batch. It’s standard in convolutional networks. During training it uses the batch’s statistics and keeps running averages; at inference it uses those averages - which is why models have separate train and eval modes.
  • Layer normalization computes μ and σ for each example across its features. It doesn’t depend on the batch at all, so it works with batch size 1 and with sequences - it’s the normalization in every transformer.
norms.py
1import numpy as np
2
3x = np.array([[1.0, 2.0, 3.0, 4.0],
4              [10.0, 10.0, 20.0, 40.0]])
5eps = 1e-5
6
7batch_norm = (x - x.mean(axis=0)) / np.sqrt(x.var(axis=0) + eps)                                   # per feature
8layer_norm = (x - x.mean(axis=1, keepdims=True)) / np.sqrt(x.var(axis=1, keepdims=True) + eps)    # per example
9print(np.round(batch_norm, 2))
10print(np.round(layer_norm, 2))
Output
[[-1. -1. -1. -1.]
 [ 1.  1.  1.  1.]]
[[-1.34 -0.45  0.45  1.34]
 [-0.82 -0.82  0.    1.63]]

Residual connections

The other big breakthrough of 2015 was the residual network (ResNet). Instead of a block computing y=F(x)y = F(x), it computes:

y=x+F(x)y = x + F(x)

The block only has to learn the change to its input, and the + x skip connection gives gradients a highway straight back through the network: the derivative of x+F(x)x + F(x) always includes a 1. ResNets trained 152 layers when 20 had been hard. Today almost every deep architecture - including every transformer - is built from residual blocks.

Key takeaways

  • Signals and gradients shrink or explode through many layers unless weights are scaled carefully.

  • Use Xavier (√(1/nᵢₙ)) for tanh/sigmoid and He (√(2/nᵢₙ)) for ReLU.

  • Batch norm normalizes features across the batch (train/eval modes differ); layer norm normalizes each example.

  • Residual connections, y = x + F(x), give gradients a direct path and make very deep networks trainable.

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

The initialization experiment

+25 XP

The input is a depth. Push a batch of 100 random inputs (width 128, rng = np.random.default_rng(0), drawn first) through that many ReLU layers, drawing each layer’s weights from the same rng, for three initializations in this order - standard deviation 0.01, then 1.0, then He (2/128\sqrt{2/128}) - restarting the rng from seed 0 for each. Print each final activation standard deviation in scientific notation with 2 decimals ({:.2e}), like std=0.01: 1.23e-45.

  • 10 layers
  • 30 layers
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

Write layer normalization

+25 XP

Implement layer_norm(x, gamma, beta, eps=1e-5): normalize each row (example) to mean 0 and variance 1 over its features, then scale by gamma and shift by beta (one value per feature). Each input line is one example. Print the normalized rows to 3 decimals, then each output row’s mean and standard deviation to 3 decimals.

  • Two examples
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: