Md. Asif Uddin
    Problem I.8.B02

    Differentiate the statistics, not just the numerator

    symbolic▲▲▲

    Exact expressions first; four decimal places unless stated otherwise.

    STATEMENT

    Derive the batch-normalisation backward pass including epsilon. Check it against finite differences, and identify the term lost by treating the variance as a constant.

    GIVEN

    A group of mm scalars, μ=m−1∑ixi\mu=m^{-1}\sum_i x_i, v=m−1∑i(xi−μ)2v=m^{-1}\sum_i(x_i-\mu)^2, r=v+εr=\sqrt{v+\varepsilon}, x^i=(xi−μ)/r\hat x_i=(x_i-\mu)/r, and yi=γx^i+βy_i=\gamma\hat x_i+\beta. The upstream derivative is gi=∂L/∂yig_i=\partial\mathcal L/\partial y_i. For a numerical check take x=(−1,0,1)x=(-1,0,1), ε=1/3\varepsilon=1/3, γ=2\gamma=2, β=0\beta=0, and g=(1,−1,2)g=(1,-1,2).

    FIND

    The full Jacobian of the standardisation, input gradients, affine-parameter gradients, and the reason the input gradients sum to zero. Decide whether the gradient is also orthogonal to the centred input when epsilon is positive.

    STRATEGY

    Differentiate the mean first, then the variance, then the inverse standard deviation. Only then contract the Jacobian with the upstream derivative.

    SOLUTION

    Write ci=xi−μc_i=x_i-\mu. Since ∑ici=0\sum_i c_i=0,

    ∂ci∂xj=δij−1m,∂v∂xj=2m∑ici(δij−1m)=2cjm.\begin{aligned} \frac{\partial c_i}{\partial x_j}&=\delta_{ij}-\frac1m,\\ \frac{\partial v}{\partial x_j} &=\frac2m\sum_i c_i\left(\delta_{ij}-\frac1m\right)\\ &=\frac{2c_j}{m}. \end{aligned}

    Consequently ∂r/∂xj=cj/(mr)\partial r/\partial x_j=c_j/(mr). The product rule applied to cir−1c_i r^{-1} gives

    Jij=δij−1/mr−cicjmr3=1r(δij−1m−x^ix^jm),\begin{aligned} J_{ij}&=\frac{\delta_{ij}-1/m}{r}-\frac{c_i c_j}{mr^3}\\ &=\frac1r\left(\delta_{ij}-\frac1m-\frac{\hat x_i\hat x_j}{m}\right), \end{aligned}

    which is (I.8.3). Omitting the derivative of the variance loses the last term.

    Contract with γgi\gamma g_i and sum over ii:

    ∂L∂xj=γr(gj−g‾−x^jgx^‾).\frac{\partial\mathcal L}{\partial x_j} =\frac\gamma r\left(g_j-\overline g-\hat x_j\overline{g\hat x}\right).

    This is (I.8.4). The parameter derivatives are

    ∂L∂γ=∑igix^i,∂L∂β=∑igi.\begin{aligned} \frac{\partial\mathcal L}{\partial\gamma}&=\sum_i g_i\hat x_i,\\ \frac{\partial\mathcal L}{\partial\beta}&=\sum_i g_i. \end{aligned}

    For the supplied check, μ=0\mu=0, v=2/3v=2/3, and r=1r=1 exactly. Thus x^=(−1,0,1)\hat x=(-1,0,1), g‾=2/3\overline g=2/3, and gx^‾=1/3\overline{g\hat x}=1/3. Substitution gives

    ∇xL=(4/3,−10/3,2),∂γL=1,∂βL=2.\begin{aligned} \nabla_x\mathcal L&=(4/3,-10/3,2),\\ \partial_\gamma\mathcal L&=1,\qquad \partial_\beta\mathcal L=2. \end{aligned}

    To check an input derivative numerically, define F(x)=∑igiyi(x)F(x)=\sum_i g_i y_i(x) and recompute all statistics after each perturbation:

    djFD=F(x+hej)−F(x−hej)2h.d_j^{\mathrm{FD}} =\frac{F(x+h e_j)-F(x-h e_j)}{2h}.

    The reproduction script uses h=10−5h=10^{-5} and asserts a maximum absolute disagreement below 10−810^{-8}. It also tests a non-symmetric input, so the check does not rely only on r=1r=1.

    Summing the analytic gradients cancels both mean terms because ∑jx^j=0\sum_j\hat x_j=0. Their sum is zero. A common shift in every input cannot change the normalised output.

    The centred input is different. Direct multiplication gives

    Jc=εr3c.J\vec c=\frac{\varepsilon}{r^3}\vec c.

    For nonzero centred data this vanishes only when epsilon is zero. In the numerical example cT∇xL=2/3\vec c^{\mathsf T}\nabla_x\mathcal L=2/3, not zero. Exact scale invariance would predict the wrong answer.

    Answer

    The Jacobian is (I.8.3), and the efficient input derivative is (I.8.4). The checked gradient is (1.333333,−3.333333,2.000000)(1.333333,-3.333333,2.000000). Its sum is zero, but its dot product with the centred input is 2/32/3.

    Check — sanity

    A constant upstream gradient produces zero input gradient when gamma is shared over the batch group: it asks to change a sum that normalisation has fixed. The computation needs two reductions and one elementwise pass, not storage of an mm-by-mm matrix.

    Where this breaks

    This is the training derivative. Frozen evaluation statistics yield only the diagonal derivative γ/rrun\gamma/r_{\mathrm{run}}. For layer normalisation, feature-specific gamma values must multiply the upstream gradient before the group reductions.

    Variation

    Derive the layer-normalisation input derivative with distinct γi\gamma_i. Set ui=γigiu_i=\gamma_i g_i and replace the expression by (uj−u‾−x^jux^‾)/r(u_j-\overline u-\hat x_j\overline{u\hat x})/r. Test why substituting the average gamma instead is generally wrong.

    Draws on