skip to content

Why does He initialization use a weight variance of 2/fan_in for ReLU layers?

level: middleimportance: must knowfreq 62%

answer

  1. write the per-layer variance multiplier
  2. fan_in times weight variance
  3. ReLU discards the negative half
  4. half the second moment must be repaid
  5. solve multiplier equals one

basics

~10 s

ReLU zeros about half of a symmetric, zero-mean pre-activation distribution, halving the signal's second moment at every layer. Doubling the weight variance from 1/fan_in to 2/fan_in cancels that halving, so activation scale survives depth.

solid answer

~50 s

Start from the forward variance recursion. For a layer with `fan_in` inputs and zero-mean independent weights, the pre-activation variance is `Var(y) = fan_in * Var(W) * E[x^2]`, where `x` is the incoming activation. To keep the scale constant across layers you want that multiplier to be one. For a near-identity activation such as tanh near zero, that gives `Var(W) = 1/fan_in`, which is Glorot's forward target -- Xavier/Glorot then averages it with the backward target `1/fan_out` and uses `2/(fan_in + fan_out)`. ReLU changes the accounting: for a symmetric zero-mean pre-activation it discards the negative half, so `E[relu(y)^2] = 0.5 * Var(y)`. Insert that factor and the condition becomes `fan_in * Var(W) / 2 = 1`, so `Var(W) = 2/fan_in`. Use the tanh-style rule on a deep ReLU stack and each layer loses half its second moment: over 20 layers that is a factor of 2^20 in variance, roughly a thousandth in magnitude.

code

python · 17 lines
python
import math, random

def final_rms(w_var, depth=20, fan_in=256, width=256):
    x = [random.gauss(0, 1) for _ in range(fan_in)]    # input RMS is 1.0
    sd = math.sqrt(w_var)
    for _ in range(depth):
        nxt = []
        for _ in range(width):
            pre = sum(random.gauss(0, sd) * xi for xi in x)
            nxt.append(pre if pre > 0 else 0.0)        # ReLU
        x = nxt
    return math.sqrt(sum(v * v for v in x) / len(x))

random.seed(0)
n = 256
print(round(final_rms(1.0 / n), 5))   # 0.00126 -- 1/fan_in, the square-layer Xavier value
print(round(final_rms(2.0 / n), 5))   # 0.81422 -- He's 2/fan_in holds the scale

go deeper

for a junior

Know the names and what they key off: Xavier/Glorot for tanh-like activations, He for ReLU, and that both set the spread from the layer's fan-in rather than from a fixed number you like the look of.

for a middle

Be able to derive it live: write the per-layer multiplier as fan_in times the weight variance, insert the one-half that ReLU costs, and solve for the variance that makes the product one.

for a senior

Expect to connect the algebra to a symptom you have actually seen -- activation magnitudes collapsing or saturating across depth at step zero -- and to say which number you change and by how much.

for a principal

Be ready to argue how much initialization tuning is worth in a codebase where rescaling layers and residual paths already stabilize most runs, and where you would still mandate an explicit per-layer scale rule as cheap insurance.

## The quantity being preserved Initialization scale is not a matter of taste; it is the solution of a small variance equation. Take one layer: `y = W x + b`, with `fan_in` incoming connections per output unit, weights drawn independently with mean zero and variance `Var(W)`, and biases at zero. `fan_in` is the number of inputs each output unit sums over; `fan_out` is the number of outputs each input feeds. Because the weights are zero-mean and independent of `x`, the pre-activation of one output unit is a sum of `fan_in` independent zero-mean terms: ``` Var(y) = fan_in * Var(W) * E[x^2] ``` Read `fan_in * Var(W)` as the **per-layer multiplier**. If it is greater than one, the signal grows geometrically with depth; if less than one, it shrinks geometrically; if exactly one, it holds. ## Glorot's answer for a near-identity activation A tanh unit is approximately the identity for small inputs, so if the signal stays small the activation barely changes the second moment: `E[x_next^2] is about Var(y)`. Setting the multiplier to one gives `Var(W) = 1/fan_in`. The backward direction gives a different answer. The signal travelling back through the same layer is summed over `fan_out` terms, so preserving *its* scale wants `Var(W) = 1/fan_out`. Unless the layer is square, no single number does both. Glorot's compromise splits the difference: ``` Var(W) = 2 / (fan_in + fan_out) ``` which for a square layer collapses back to `1/fan_in`. The uniform version draws from `-sqrt(6/(fan_in+fan_out))` to `+sqrt(6/(fan_in+fan_out))`, since a uniform distribution on `[-a, a]` has variance `a^2/3`. ## Where the factor of two comes from ReLU is not near-identity. It passes positive inputs unchanged and maps negatives to zero. If the pre-activation `y` is symmetric around zero -- which it is, since the weights are symmetric and zero-mean -- then exactly half its probability mass is discarded and the surviving half contributes its usual share: ``` E[relu(y)^2] = 0.5 * E[y^2] = 0.5 * Var(y) ``` So the recursion across a ReLU layer is `Var(y_next) = fan_in * Var(W) * 0.5 * Var(y)`. Setting that multiplier to one: ``` fan_in * Var(W) / 2 = 1 -> Var(W) = 2 / fan_in ``` That is the whole derivation, and the factor of two is nothing more exotic than "ReLU throws away half the second moment, so put twice as much in". Note it is a *forward* rule keyed to `fan_in`; the backward variant keyed to `fan_out` exists and, in practice, either works, because the discrepancy is a factor of the width ratio, not of depth. ## What getting it wrong costs, quantitatively Take a 20-layer ReLU stack for seismic-waveform regression, width 256, initialized with the tanh-style `Var(W) = 1/fan_in`. The per-layer multiplier is `256 * (1/256) * 0.5 = 0.5`. Each layer halves the second moment, so at layer 20 the variance is `2^-20` of the input's and the typical activation magnitude is about `2^-10`, near one thousandth. The last layers are computing on numerical dust, the output is nearly constant across examples, and the run looks "slow" rather than broken. Swap to `2/fan_in` and the multiplier is `256 * (2/256) * 0.5 = 1`: the scale is flat all the way down. The symmetric failure is a fixed standard deviation. Consider a 30-layer plain tanh MLP over 60-channel machine-vibration sensor windows, hidden width 5,000, every weight drawn at a fixed 0.01 standard deviation because that "looks small". The multiplier is `5000 * (0.01)^2 = 0.5` -- activation variance halves at every layer, and by layer 30 the signal is gone. Widen the layer to 20,000 with the same 0.01 and the multiplier becomes 2 and the signal explodes into tanh saturation instead. A fixed standard deviation is not a scale rule at all: it silently means something different at every width, which is exactly why the rules are stated in terms of fan-in. ## The assumptions, and where they leak The derivation assumes zero-mean independent weights, independence between weights and inputs, and a symmetric pre-activation distribution. These are true at step zero and progressively less true afterwards -- correlations build up as training proceeds. That is fine: the rule is a rule for *step zero*, whose job is to get the run onto a reasonable trajectory, not to hold for all time. One more consequence worth internalising: the rules are about **variance**, so the choice of Gaussian versus uniform is essentially cosmetic. What matters is the second moment and the fact that the draw is random across the width. ## Why the choice feels less critical in modern networks Architectures that rescale each layer's output absorb a constant multiplicative error in the incoming weights, so a mis-scaled initialization that would have killed a plain stack often merely costs some early progress. The scale still matters where such layers are absent or infrequent, in very deep stacks where per-layer errors compound before the first rescaling, and because the weight magnitude interacts with how large a *relative* step each update makes. Treating an explicit fan-in-based rule as free insurance is the right default -- it costs one line and removes an entire failure class.

  • What does Xavier/Glorot use instead, and why does the fan_out term appear?
    Glorot targets both directions at once. Preserving the forward activation scale wants `Var(W) = 1/fan_in`; preserving the scale of the signal travelling backwards wants `1/fan_out`. No single value does both unless the layer is square, so the compromise is `Var(W) = 2/(fan_in + fan_out)`. It assumes an activation that is near-identity around zero, such as tanh in its linear region, which is precisely why ReLU needs the separate factor of two.
  • How should the scale for a 50,000-row categorical embedding table differ from the dense layer it feeds?
    An embedding row is selected by a lookup rather than summed over 50,000 inputs, so the table height is not a fan-in in the variance argument -- each output is one row, not a sum. Applying a `1/sqrt(50000)` style scale makes the vectors vanishingly small and starves the layer above. Pick the scale from the activation scale you want the embedding vector to have, or from the fan-out of the dimension it feeds, not from the number of rows.
  • Does the rule change if the layer's inputs are not zero-mean?
    Yes, in the sense that the derivation's clean form breaks. The recursion uses `E[x^2]`, not `Var(x)`, and for a mean-shifted input those differ by the square of the mean, so the effective multiplier is larger than the formula suggests and the signal grows faster than intended. This is one reason raw features are usually centred before the first layer, and why post-ReLU activations -- which are strictly nonnegative -- must be tracked through their second moment rather than their variance.

saying these in an interview costs you the question

  • Says the factor of two comes from ReLU having two branches
  • Applies the tanh-style scale to a deep ReLU stack unchanged
  • Claims a fixed 0.01 standard deviation is safe at any width
  • Thinks Gaussian versus uniform matters more than the variance
  • Confuses picking an initialization scale with rescaling activations at run time

context