Variance propagation and initialization
Overview
Section titled “Overview”| Module | M07.3 · build · Python · Pass 2 · 2 to 3 h |
| You build | python/tinyllm/nn/init.py: fans, calculate_gain, xavier_uniform, xavier_normal, kaiming_normal, normal_init, scaled_residual_std |
| Contract | course/contracts/py/tinyllm/nn/init.pyi · draw order: spec/pcg32.md |
| Tests | course/tests/M07.3/test_init.py (what they check: section 4); torch’s gain table and fan rule in course/fixtures/M07.3/init_torch.json |
| Needs | M07.0 its normal(rng, n) draws every normal weight (or --ref-deps). Reading: M01.3 (activations), M06.3 (PCG32), S-M02 and S-M04 (the Gaussian integrals) |
| Used by | L0.4 initializes every Linear and Embedding with it · later L3.2 (LSTM), L7.9 (the modern decoder), C1 · later: L3.3, L3.6, L5.4, L5.5, L6.1, L6.2, L6.3 |
| Milestone | MS-P2 (Pass 2 closes with every math module it teaches passing) |
| Optional depth | Glorot and Bengio, “Understanding the difficulty of training deep feedforward neural networks” (2010), section 4.2; He et al., “Delving Deep into Rectifiers” (2015), section 2.2; Radford et al., “Language Models are Unsupervised Multitask Learners” (GPT-2, 2019), section 2.3 |
Key Takeaways
Section titled “Key Takeaways”- For with independent zero-mean weights, : the variance of a layer’s output is set by its fan-in and the weights’ variance (
test_hand_example_linear_4x3). - Xavier picks for layers that are linear near 0; Kaiming picks for ReLU, which zeroes half its inputs (
test_empirical_variance_matches_the_formula). - With Kaiming the signal survives 20 ReLU layers; with Xavier it shrinks by about (
test_relu_signal_survives_20_layers). - A gain folds the activation’s effect into one factor; fans of a convolution or a Linear weight follow torch’s
(out, in, *kernel)rule (test_fans_and_gains_match_torch). - Weights are drawn in row-major order, one draw per element, from your PCG32, so a seed fixes the model in every language (
test_normal_inits_draw_spec_normals_in_c_order).
How to work this chapter
Section titled “How to work this chapter”ol start M07.3 # stubs python/tinyllm/nn/init.py, contract alongsideol tests M07.3 # read the test catalog first: rung R0, you write no tests hereol check M07.3 # exit code is the verdictol check M07.3 --ref-deps # only if your M07.0 is not passing yetol diff M07.3 # after passing: your code against the reference1. Why now
Section titled “1. Why now”L0.4 builds layers next, and a layer needs starting weights. The obvious choices fail. All zeros makes every unit in a layer compute the same thing and receive the same gradient forever. Standard normal weights make a layer with 512 inputs multiply the size of its input by about , so after a few layers the activations overflow float32 or saturate every tanh, and the gradients are 0 or NaN. Weights a little too small do the opposite: the signal shrinks geometrically and the deep layers learn nothing. Your MS-L0 milestone trains an MLP on digits from your own initialization, and C1 trains a 10M-parameter Transformer the same way. This module derives the one number that decides which of these happens, the variance of the weights, from expectation and variance (M07.0), and builds the initializers every later model calls.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| , | expectation and variance of a random variable, (M07.0) | scalars |
| a Linear weight, one row per output (torch’s layout) | float32[out, in] | |
| the layer’s input; its pre-activation | vectors | |
| fan-in and fan-out: inputs feeding one output, outputs fed by one input | int | |
| the variance every weight is drawn with | float | |
| the activation, | function | |
| the gain of : the factor that corrects the variance for it | float | |
| , | uniform on ; normal with mean 0 and standard deviation | distributions |
| the number of Transformer blocks | int |
2.1 Expectation and variance of sums and products
Section titled “2.1 Expectation and variance of sums and products”Two rules from M07.0 do all the work. Expectation is linear: for any random variables. Variance adds for independent ones: , because the cross term is a product of two zero means. For a product of independent and with :
Note , not : the input’s mean counts too (after a ReLU, has a positive mean).
2.2 Choosing the weight variance
Section titled “2.2 Choosing the weight variance”One output of a layer is , a sum of independent products, each with mean 0. By 2.1:
To keep the signal the same size from layer to layer, set this equal to the previous layer’s value. Linear activations (, zero mean): and the condition is . The backward pass multiplies gradients by , whose rows have entries, so it wants . Xavier (Glorot and Bengio) splits the difference:
ReLU, , keeps the positive half of a symmetric and zeroes the rest, so . The next layer then sees , and preserving the variance needs twice as much weight variance. Kaiming (He et al.):
Getting the factor wrong compounds: a variance too small by 2 per layer is after 20 layers, and too large by 2 is .
2.3 Fans and gains
Section titled “2.3 Fans and gains”A gain writes the activation’s correction as a factor on the standard deviation: Kaiming is with for ReLU, and Xavier takes an optional gain, . torch’s table, which you match: for linear, convolutions, and sigmoid; for ReLU; for leaky ReLU with negative slope (default ), because it keeps of the negative half’s second moment; for SELU; and for tanh. The tanh value is empirical: tanh has slope 1 at 0 but shrinks larger inputs, and with the mean square of a deep tanh stack settles near 0.42 instead of draining toward 0 (about 0.02 after 20 layers with ).
Fans come from the weight’s shape (out, in, *kernel): , , with the product of the kernel dimensions (1 for a Linear weight). A convolution with 4 input channels and a kernel has . Fewer than 2 dimensions have no fan-in and are an error.
2.4 Drawing the weights
Section titled “2.4 Drawing the weights”Two distributions with the same variance work equally well at the start. A uniform has variance , so Xavier-uniform uses ; element is with the -th rng.uniform(). A normal is times a standard normal from your M07.0 normal(rng, n). Either way the array is filled in row-major (C) order, one draw per element, and nothing else is drawn, so a model initialized layer after layer from one generator is reproducible, and the Rust port (L10.1) can rebuild the same weights from the same seed. Results are float32, the dtype of every parameter.
2.5 Residual streams
Section titled “2.5 Residual streams”A Transformer block adds its outputs to a running residual stream: , two additions per block, in all. If each addition is independent with variance , the stream’s variance grows to . GPT-2 scales the standard deviation of the two output projections in every block to
whatever the depth (scaled_residual_std, with in GPT-2 and your L7.9).
3. Worked example by hand
Section titled “3. Worked example by hand”A Linear layer with 3 inputs and 4 outputs stores as shape : , .
| Initializer | Formula | Value |
|---|---|---|
| Xavier-uniform bound | ||
| Xavier-normal std | ||
| Kaiming-normal std (ReLU, fan_in) | ||
| GPT-2 residual std, |
The variance carries through. Feed the Kaiming layer an input with : . ReLU keeps half the second moment: , exactly what the next layer received. With Xavier () the same input gives and : the signal loses more than half per layer.
The draws. With PCG32(0) the first uniform is , so Xavier-uniform’s first weight is . Kaiming’s first weight uses the first pair : , , so the standard normal is and the weight . The first test, test_hand_example_linear_4x3, checks every one of these against the frozen generator.
4. The interface
Section titled “4. The interface”def fans(shape: tuple[int, ...]) -> tuple[int, int] # (fan_in, fan_out)def calculate_gain(nonlinearity: str, param: float | None = None) -> floatdef xavier_uniform(shape, gain: float, rng) -> NDArray # float32, U(-a, a)def xavier_normal(shape, gain: float, rng) -> NDArray # float32, N(0, s^2)def kaiming_normal(shape, fan_mode: Literal["fan_in", "fan_out"], nonlinearity: str, rng) -> NDArraydef normal_init(shape, std: float, rng) -> NDArray # any shape, including 1-Ddef scaled_residual_std(base_std: float, n_layers: int) -> float # base / sqrt(2 L)rng is a PCG32: your M06.3 generator in your own code, the frozen course/tests/_lib/pcg32.py in the course tests. Uniforms come from rng.uniform(), normals from tinyllm.prob.rv.normal(rng, n). Every function raises ValueError for a shape without a fan (fewer than 2 dimensions, or a 0 fan), a negative gain or standard deviation, an unknown nonlinearity or fan mode, and n_layers < 1.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example_linear_4x3 | unit | section 3: fans, gains, the residual std, and the first weights of PCG32(0) | you and the tests agree on every formula |
test_fans_and_gains_match_torch | golden | torch’s gain table and fan rule for Linear and conv shapes | L0.4 compares layers with torch |
test_normal_inits_draw_spec_normals_in_c_order | unit | each normal initializer is std times the spec normals, row-major, odd sizes too | reproducible weights across languages |
test_uniform_uses_one_draw_per_element | unit | exact values, and the generator advanced by exactly one uniform per element | layers initialized in sequence stay aligned |
test_same_seed_same_weights | property | equal seeds give equal bytes, different seeds differ | replayable runs |
test_empirical_variance_matches_the_formula | statistical | mean 0 and variance within 4 standard errors for all four variants | the formulas of section 2.2, measured |
test_relu_signal_survives_20_layers | property | Kaiming keeps the mean square within a factor 16; Xavier collapses it below | deep MLPs and Transformers train at all |
test_tanh_signal_survives_with_xavier | property | gain keeps a 20-layer tanh stack’s mean square in ; gain 1 ends at less than half that | tanh and gated recurrent layers (L3.2) |
test_scaled_residual_keeps_the_stream_bounded | property | for several depths | GPT-2 style output projections (L7.9) |
test_shapes_and_dtype | unit | float32, C-contiguous, exact shapes, 1-D for normal_init | parameters are float32 (L0.4) |
test_rejects_bad_arguments | boundary | 1-D fans, unknown names, negative std or depth raise ValueError | no silent wrong-scale default |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
reading a (out, in) weight as (in, out) | fan-in and fan-out swap; non-square layers get the wrong scale | test_hand_example_linear_4x3, test_fans_and_gains_match_torch (mutant s01) |
| ignoring the kernel dimensions | a convolution is initialized 9 times too large in variance | test_fans_and_gains_match_torch (mutant s02) |
| using as the uniform bound | the uniform’s variance is , so this is 3 times too small | test_empirical_variance_matches_the_formula (mutant s03) |
| ReLU gain 2 (a variance) instead of (a standard deviation) | activations double in mean square per layer: after 20 | test_relu_signal_survives_20_layers (mutant s05) |
| GPT-2 residual scale instead of | each block adds two outputs; the stream ends twice as large | test_scaled_residual_keeps_the_stream_bounded (mutant s08) |
| filling in column-major order | same distribution, different weights: Python and the Rust port disagree for the same seed | test_normal_inits_draw_spec_normals_in_c_order (mutant s09) |
| making a new generator inside the initializer | every layer gets the same weights, and the seed does nothing | test_same_seed_same_weights (mutant s13) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Forward | L3.3 | Registered call site uses this module. |
| Forward | L3.6 | Registered call site uses this module. |
| Forward | L5.4 | Registered call site uses this module. |
| Forward | L5.5 | Registered call site uses this module. |
| Forward | L6.1 | Registered call site uses this module. |
| Forward | L6.2 | Registered call site uses this module. |
| Forward | L6.3 | Registered call site uses this module. |
| Direction | Module | How it uses this |
|---|---|---|
| Back | M07.0 | normal(rng, n) draws every normal weight; expectation and variance are its definitions |
| Back | M01.3 | the activations whose second moments set the gains (reading) |
| Back | M06.3 | the PCG32 generator passed in as rng (reading) |
| Back | S-M02, S-M04 | the Gaussian integrals behind (reading) |
| Forward | L0.4 | Linear and Embedding initialize their parameters with these functions |
| Forward | L3.2 | LSTM gate weights |
| Forward | L7.9 and C1 | normal_init(0.02) plus scaled_residual_std for the output projections of a Llama-style decoder |
If you skip this module, ol check L0.4 stops with BLOCKED ... needs M07.3: build it, or pass --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
xavier_*, kaiming_normal | torch.nn.init | in-place initializers, trunc_normal_, orthogonal_ (your M03.3), kaiming_uniform_ with a = sqrt(5) (torch’s Linear default) | torch/nn/init.py |
scaled_residual_std | GPT-NeoX, Megatron scaled_init_method_normal | the same rule, applied by parameter name | megatron/core/utils.py |
| a fixed base std | muP (maximal update parametrization) | width-dependent init and learning rates so hyperparameters transfer from small to large models | Yang et al., “Tensor Programs V” (2022) |
| variance analysis at initialization | signal propagation theory | mean-field analysis of depth and the edge of chaos; Fixup and T-Fixup train without normalization | Schoenholz et al., “Deep Information Propagation” (2017) |