In a 200-block residual network, why zero-initialize each block's final scale parameter?
answer
- what does a fresh 200-deep stack output?
- each branch adds its own variance
- make every block a pass-through
- one parameter zeroed, the rest random
- the gate still gets a gradient
basics
~20 sIt makes every residual branch output exactly zero at step zero, so the stack starts as the identity map and the signal keeps the scale it entered with instead of growing block by block. The branches then switch on gradually.
solid answer
~50 sA residual block computes `out = x + branch(x)`. If the branch is live at initialization, it adds its own variance to the residual stream, so the stream's scale grows with the number of blocks -- at 200 blocks the output is far out of scale before a single gradient step, and the first updates go on fixing that rather than learning. Zeroing the last learnable multiplicative scale in each branch sets `branch(x) = 0` exactly, making every block the identity and the whole stack a pass-through. The gradient with respect to that scale is still nonzero -- it is the incoming gradient dotted with the branch's pre-scale output -- so the scale leaves zero on the first step, and the inner weights, whose gradients are proportional to it, begin learning shortly after. Blocks turn themselves on in rough order of usefulness.
go deeper
Recall that a residual block adds its branch output to its input, and that starting the branch's output at zero makes the block do nothing at first, so the network begins as a pass-through.
Be able to explain why the residual stream's scale grows with block count when every branch is live at step zero, since independent contributions add in variance.
Expect to justify it operationally: name the single parameter you zero, show the derivative that proves the branch still wakes up, and describe the input-versus-output check that confirms the identity start.
Own the tradeoff between architectural fixes that make depth trainable by construction and run-time fixes such as a cautious opening step size, and say which of the two you would standardize across a team's models and why.
## What goes wrong at 200 blocks A residual block adds a computed correction to its input: `out = x + branch(x)`. The input path is an exact copy, so the block's output is at least as large as its input. If the branch is randomly initialized and active, it contributes a roughly independent term, and independent contributions add in variance: ``` Var(out) = Var(x) + Var(branch(x)) ``` Stack that 200 times and the residual stream's variance grows monotonically with depth -- linearly if each branch emits an order-one signal regardless of its input, geometrically if each branch's output scales with its input. Either way the network's output at step zero is nothing like the scale the loss expects, the first steps are spent undoing the accumulated growth rather than learning anything, and the very deep stack is fragile precisely where it was supposed to be strong. Note what is *not* the fix: shrinking the per-layer initialization inside the branch. The problem is not one branch being too big; it is 200 correct-sized branches summing. You want the depth-dependent term controlled, and per-layer fan-in rules know nothing about depth. ## The zeroed scale Most residual branches end in a multiplicative learnable parameter -- a per-channel scale, or a single learnable scalar placed there for the purpose. Initialize *that one parameter* to zero and the branch's output is identically zero for every input. Then: - `out = x + 0 = x`, so every block is exactly the identity; - the 200-block stack passes its input straight through to the head; - the network's output scale at step zero equals the scale set by the stem, independent of depth. The network begins life as a shallow network -- a stem and a head -- and grows deeper as training proceeds. That is a much better-behaved starting point than a 200-deep random function. ## Why the branch is not permanently dead This is the objection worth being ready for, and the answer is a one-line derivative. Write the branch output as `g * u(x)` where `g` is the zeroed scale and `u` is everything before it. The gradient with respect to `g` is the incoming gradient dotted with `u(x)`, and `u` is *randomly* initialized, so `u(x)` is nonzero and generically not orthogonal to the incoming gradient. `g` therefore moves off zero on the very first step. The gradient with respect to the inner weights is proportional to `g`, so those weights are frozen at step zero and start moving as soon as `g` becomes nonzero. The dynamics that follows is the point of the trick: a block only invests in its inner weights once its own gate has decided the block is worth something. Blocks turn on in rough order of usefulness rather than all shouting at once. This also explains why zeroing the *whole* branch is a mistake rather than a stronger version of the same idea. If `u`'s weights are zero too, `u(x)` is zero, the gradient with respect to `g` is zero, and the branch really is dead -- and you have reintroduced the symmetry problem inside the branch on top of it. Exactly one parameter is zeroed; everything before it stays random. ## Alternatives when there is no scale parameter to zero If the branch has no trailing multiplicative parameter, the same effect is available by construction: - **Insert a learnable scalar** at the end of each branch, initialized at zero. Cheap: one parameter per block. - **Divide each branch's output by a fixed `sqrt(L)`**, where `L` is the number of blocks. Summing `L` independent contributions each scaled by `1/sqrt(L)` keeps the total added variance order one no matter how deep the stack. This does not give an exact identity start, but it makes the stream's scale depth-independent, which was the underlying goal. - **Downscale the branch's own last weight matrix** by a depth-dependent factor. Same idea, less explicit, harder to audit later. ## When it does not apply The trick works *only* because an identity path carries the signal past the zeroed branch. Zero a layer's output scale in a plain stack with no skip connection and you have zeroed the forward signal entirely: everything downstream sees zero, the gradients arriving from below are zero, and nothing recovers. The skip is what makes a switched-off block harmless rather than fatal, so "zero the last scale" is a residual-architecture technique and not a general initialization rule. ## How you verify it One check, before training: run a batch through the untrained network and compare the tensor entering the first block with the tensor leaving the last. Under a correct zeroed-scale initialization they are identical to numerical precision. If they are not, some branch has a live path you did not account for -- an unzeroed scale, a bias after the scale, or a branch whose output is not gated by the parameter you zeroed. That check takes a minute and catches the whole class of mistakes here.
- If the scale is zero, how do the branch's inner weights ever receive a gradient?Write the branch as `g * u(x)` with `g` the zeroed scale. The gradient with respect to `g` is the incoming gradient dotted with `u(x)`, and `u` is randomly initialized, so that product is nonzero and `g` leaves zero on the first step. The inner weights' gradients are proportional to `g`, so they are frozen for one step and then start moving. Blocks effectively switch themselves on.
- Would zeroing the entire branch, not just its final scale, work even better?No, it breaks the mechanism. With every weight in the branch zeroed, the pre-scale output is zero, so the gradient with respect to the scale parameter is zero too and the branch never wakes up. You would also have recreated the symmetry problem inside the branch, since all its units would be identical. Exactly one parameter is zeroed and everything before it stays randomly initialized.
- What would you do if the architecture has no trailing scale parameter to zero?Add a single learnable scalar at the end of each branch and initialize it to zero -- one parameter per block, and the identity start is exact. If touching the architecture is off the table, multiply each branch output by a fixed `1/sqrt(L)` for a stack of `L` blocks; summing `L` independent contributions at that scale keeps the total added variance order one, so the residual stream's scale no longer grows with depth.
- Does the same trick help in a plain 200-layer stack with no skip connections?No, and applying it there is actively destructive. Without an identity path, zeroing a layer's output scale zeroes the forward signal for everything downstream, and the gradient flowing back to the layers below is zero as well -- the network computes a constant and cannot escape. The technique depends entirely on the skip carrying the signal past a switched-off branch.
saying these in an interview costs you the question
- Says a zeroed scale means the block can never learn
- Zeroes the whole branch rather than its final scale
- Believes it replaces the fan-in scale rule inside the branch
- Claims it works the same without a skip connection
- Thinks the problem is one oversized branch rather than 200 summed ones