Practice / Minibatches

Batch norm backward

Ten problems on the backward pass of batch normalisation: the gradients of γ, β and the batch statistics, the compact input gradient and why each feature's column of it sums to zero, batch norm as layer norm on the transpose, inference with running statistics, a batch of one and the role of ε, with worked solutions and the mistakes that hold the batch statistics constant or take them over the wrong axis.

Before you start

Batch normalisation is the one common layer in which the examples of a minibatch interact. It standardises each feature using the mean and variance of that feature over the batch, then rescales and shifts with learned γ\gamma and β\beta. Because the statistics are computed from the batch, every example's output depends on every other example's input, and the backward pass has to follow those paths. These ten problems derive the gradient of every piece, assemble the compact input gradient that most implementations use, compare it with layer norm, and then look at the cases where it changes character: inference with running statistics, a batch of one, and a feature that is constant over the batch. The five mistakes at the end are the ones that still produce a plausible gradient: the statistics taken over the wrong axis, the batch statistics held constant, an unbiased variance differentiated against a biased forward pass, γ\gamma's gradient left unsummed, and the running variance used where the batch variance belongs.

  • The conventions are those of the previous pages: a gradient has the shape of the variable it is taken with respect to, ⊙\odot is the elementwise product, 1\mathbf{1} is the all-ones vector, and rows are examples. As on the previous page, dAdA is the code-style name for ∇AL\nabla_A L, with the shape of AA.
  • The input is X∈RN×dX \in \mathbb{R}^{N\times d}: NN examples, dd features. Batch norm works down the columns. Feature jj has batch mean μj=1N∑nXnj\mu_j = \tfrac1N\sum_n X_{nj} and biased variance σj2=1N∑n(Xnj−μj)2\sigma_j^2 = \tfrac1N\sum_n (X_{nj} - \mu_j)^2, and μ,σ2∈Rd\mu, \sigma^2 \in \mathbb{R}^d collect them.
  • ϵ>0\epsilon > 0 is a small constant inside the square root. DD is the d×dd\times d diagonal matrix with Djj=(σj2+ϵ)−1/2D_{jj} = (\sigma_j^2+\epsilon)^{-1/2}, so the normalised input is X^=(X−1μ⊤)D\hat X = (X - \mathbf{1}\mu^\top)D, with entries X^nj=(Xnj−μj)/σj2+ϵ\hat X_{nj} = (X_{nj} - \mu_j)/\sqrt{\sigma_j^2+\epsilon}.
  • The output is Y=X^diag⁡(γ)+1β⊤Y = \hat X\operatorname{diag}(\gamma) + \mathbf{1}\beta^\top, that is Ynj=γjX^nj+βjY_{nj} = \gamma_j\hat X_{nj} + \beta_j, with learned γ,β∈Rd\gamma, \beta \in \mathbb{R}^d. LL is a scalar loss that depends on XX, γ\gamma and β\beta only through YY, and dYdY is given.
  • One feature at a time: x∈RNx \in \mathbb{R}^N is column jj of XX, μ\mu and σ2\sigma^2 are then the scalars μj\mu_j and σj2\sigma_j^2, x^=(x−μ1)/σ2+ϵ\hat x = (x - \mu\mathbf{1})/\sqrt{\sigma^2+\epsilon} is column jj of X^\hat X, and dx^d\hat x, dxdx are column jj of dX^d\hat X, dXdX. Sums ∑n\sum_n run over the NN examples.
  • Running statistics μˉ,σˉ2∈Rd\bar\mu, \bar\sigma^2 \in \mathbb{R}^d are moving averages of the batch statistics, updated outside the gradient; at inference they replace μ\mu and σ2\sigma^2, and Dˉ\bar D is DD built from σˉ2\bar\sigma^2.

Builds on: Batched backprop: dense layers on a minibatch, Layer norm backward

Problems

  1. ·

    Give the shapes of μ\mu, σ2\sigma^2, X^\hat X, γ\gamma, β\beta and YY. Which entries of XX does YnjY_{nj} depend on?

  2. ·

    Compute dγd\gamma, dβd\beta and dX^d\hat X from dYdY.

  3. ··

    For one feature x∈RNx \in \mathbb{R}^N, compute ∇xμ\nabla_x\mu and ∇xσ2\nabla_x\sigma^2.

  4. ··

    Treat one feature's forward pass as a graph: μ\mu is computed from xx, σ2=1N∑n(xn−μ)2\sigma^2 = \tfrac1N\sum_n (x_n-\mu)^2 from xx and μ\mu, and x^n=(xn−μ)/σ2+ϵ\hat x_n = (x_n-\mu)/\sqrt{\sigma^2+\epsilon} from xx, μ\mu and σ2\sigma^2. Given dx^d\hat x, compute the gradients dσ2=∂L/∂σ2d\sigma^2 = \partial L/\partial\sigma^2 and dμ=∂L/∂μd\mu = \partial L/\partial\mu arriving at those two nodes.

  5. ···

    Assemble dxdx for one feature from dx^d\hat x, dσ2d\sigma^2 and dμd\mu, and simplify it to a form that uses only dx^d\hat x, x^\hat x and σ2\sigma^2.

  6. ··

    For one feature, show that 1⊤dx=0\mathbf{1}^\top dx = 0 and compute x^⊤dx\hat x^\top dx. When is dxdx orthogonal to x^\hat x?

  7. ··

    Let LN⁡(Z)\operatorname{LN}(Z) be layer norm without γ\gamma and β\beta, applied to each row of a matrix ZZ with that row's own mean, biased variance and ϵ\epsilon. Express X^\hat X and dXdX through LN⁡\operatorname{LN}, and say which axis each layer averages over.

  8. ··

    At inference the layer uses the running statistics: X^=(X−1μˉ⊤)Dˉ\hat X = (X - \mathbf{1}\bar\mu^\top)\bar D. Compute dXdX, dγd\gamma and dβd\beta.

  9. ·

    Train with a batch of one (N=1N = 1). Compute μ\mu, σ2\sigma^2, X^\hat X, YY and the gradients dXdX, dγd\gamma, dβd\beta.

  10. ···

    With N≥2N \ge 2, one feature takes the same value on every example in the batch. Compute x^\hat x and dxdx for that feature. What does ϵ\epsilon do here, and how large is the effect?

Worked solutions

Problem 1

Give the shapes of μ\mu, σ2\sigma^2, X^\hat X, γ\gamma, β\beta and YY. Which entries of XX does YnjY_{nj} depend on?

  1. μ=1NX⊤1∈Rd\mu = \tfrac1N X^\top\mathbf{1} \in \mathbb{R}^d, and likewise σ2∈Rd\sigma^2 \in \mathbb{R}^d.Each is a sum over the NN examples, which leaves one number per feature; (d×N)(N×1)=d×1(d\times N)(N\times 1) = d\times 1.
  2. X−1μ⊤X - \mathbf{1}\mu^\top is N×dN\times d and subtracts μj\mu_j from every entry of column jj; multiplying on the right by the diagonal DD scales column jj by (σj2+ϵ)−1/2(\sigma_j^2+\epsilon)^{-1/2}, so X^\hat X is N×dN\times d.(1μ⊤)nj=μj(\mathbf{1}\mu^\top)_{nj} = \mu_j, and (AD)nj=AnjDjj(AD)_{nj} = A_{nj}D_{jj} for a diagonal DD.
  3. Y=X^diag⁡(γ)+1β⊤Y = \hat X\operatorname{diag}(\gamma) + \mathbf{1}\beta^\top is N×dN\times d, with γ,β∈Rd\gamma, \beta \in \mathbb{R}^d.One scale and one shift per feature, shared by every example, as the bias of the previous page's dense layer is.
  4. X^nj\hat X_{nj} uses XnjX_{nj}, μj\mu_j and σj2\sigma_j^2, and μj\mu_j and σj2\sigma_j^2 are built from all NN entries of column jj and nothing else.The definitions of μj\mu_j and σj2\sigma_j^2 sum over nn with jj fixed.
  5. μ,σ2,γ,β∈Rd\mu, \sigma^2, \gamma, \beta \in \mathbb{R}^d; X^,Y∈RN×d\hat X, Y \in \mathbb{R}^{N\times d}; YnjY_{nj} depends on every entry of column jj of XX and on no other columnFeatures never mix and examples always do, the reverse of the dense layer on the previous page, where row nn of the output used only row nn of the input. So everything below can be done one column at a time, and this is the layer the previous page's Problem 9 excluded from gradient accumulation.

Problem 2

Compute dγd\gamma, dβd\beta and dX^d\hat X from dYdY.

  1. Ynj=γjX^nj+βjY_{nj} = \gamma_j\hat X_{nj} + \beta_j.Index form of Y=X^diag⁡(γ)+1β⊤Y = \hat X\operatorname{diag}(\gamma) + \mathbf{1}\beta^\top; X^\hat X does not depend on γ\gamma or β\beta.
  2. ∂L/∂γj=∑ndYnjX^nj\partial L/\partial\gamma_j = \sum_n dY_{nj}\hat X_{nj} and ∂L/∂βj=∑ndYnj\partial L/\partial\beta_j = \sum_n dY_{nj}.γj\gamma_j and βj\beta_j appear in YnjY_{nj} for every example nn, with coefficients X^nj\hat X_{nj} and 11, and the chain rule sums over every entry that contains them.
  3. ∂L/∂X^nj=γj dYnj\partial L/\partial\hat X_{nj} = \gamma_j\,dY_{nj}.X^nj\hat X_{nj} appears only in YnjY_{nj}, with coefficient γj\gamma_j.
  4. For a matrix MM with NN rows, M⊤1M^\top\mathbf{1} adds up its rows, and Mdiag⁡(γ)M\operatorname{diag}(\gamma) scales column jj by γj\gamma_j.(M⊤1)j=∑nMnj(M^\top\mathbf{1})_j = \sum_n M_{nj} and (Mdiag⁡(γ))nj=Mnjγj(M\operatorname{diag}(\gamma))_{nj} = M_{nj}\gamma_j.
  5. dγ=(dY⊙X^)⊤1d\gamma = (dY\odot\hat X)^\top\mathbf{1}; dβ=dY⊤1d\beta = dY^\top\mathbf{1}; dX^=dYdiag⁡(γ)d\hat X = dY\operatorname{diag}(\gamma)Shapes dd, dd and N×dN\times d. The sums over the batch appear because γ\gamma and β\beta are shared by every example. From here on the input gradient needs only dX^d\hat X, one column at a time: for feature jj, dx^=γjd\hat x = \gamma_j times column jj of dYdY.

Problem 3

For one feature x∈RNx \in \mathbb{R}^N, compute ∇xμ\nabla_x\mu and ∇xσ2\nabla_x\sigma^2.

  1. ∂μ/∂xm=1N\partial\mu/\partial x_m = \tfrac1N for every mm.μ=1N∑nxn\mu = \tfrac1N\sum_n x_n contains xmx_m once, with coefficient 1N\tfrac1N.
  2. ∂σ2/∂xm=2N∑n(xn−μ)(∂xn/∂xm−∂μ/∂xm)\partial\sigma^2/\partial x_m = \tfrac2N\sum_n (x_n-\mu)\big(\partial x_n/\partial x_m - \partial\mu/\partial x_m\big).Chain rule on each square (xn−μ)2(x_n-\mu)^2, with μ\mu kept as the function of xx that it is.
  3. ∂σ2/∂xm=2N(xm−μ)−2N2∑n(xn−μ)\partial\sigma^2/\partial x_m = \tfrac2N(x_m-\mu) - \tfrac2{N^2}\sum_n (x_n-\mu).∂xn/∂xm\partial x_n/\partial x_m is 11 for n=mn = m and 00 otherwise, which picks out one term; step 1 gives ∂μ/∂xm=1N\partial\mu/\partial x_m = \tfrac1N, the same for every nn.
  4. ∑n(xn−μ)=∑nxn−Nμ=0\sum_n (x_n - \mu) = \sum_n x_n - N\mu = 0.Nμ=∑nxnN\mu = \sum_n x_n by the definition of μ\mu.
  5. ∇xμ=1N1\nabla_x\mu = \tfrac1N\mathbf{1}; ∇xσ2=2N(x−μ1)\nabla_x\sigma^2 = \tfrac2N(x - \mu\mathbf{1})Both N×1N\times 1, the shape of xx. The 1N\tfrac1N in ∇xσ2\nabla_x\sigma^2 is there because the forward pass uses the biased variance; it is the layer-norm page's Problem 2 with the batch in place of the features.

Problem 4

Treat one feature's forward pass as a graph: μ\mu is computed from xx, σ2=1N∑n(xn−μ)2\sigma^2 = \tfrac1N\sum_n (x_n-\mu)^2 from xx and μ\mu, and x^n=(xn−μ)/σ2+ϵ\hat x_n = (x_n-\mu)/\sqrt{\sigma^2+\epsilon} from xx, μ\mu and σ2\sigma^2. Given dx^d\hat x, compute the gradients dσ2=∂L/∂σ2d\sigma^2 = \partial L/\partial\sigma^2 and dμ=∂L/∂μd\mu = \partial L/\partial\mu arriving at those two nodes.

  1. σ2\sigma^2 feeds only the x^n\hat x_n, and ∂x^n/∂σ2=−12(xn−μ)(σ2+ϵ)−3/2\partial\hat x_n/\partial\sigma^2 = -\tfrac12(x_n-\mu)(\sigma^2+\epsilon)^{-3/2}.In the graph each node's inputs are held fixed when it is varied; power rule on (σ2+ϵ)−1/2(\sigma^2+\epsilon)^{-1/2}.
  2. dσ2=∑ndx^n⋅(−12)(xn−μ)(σ2+ϵ)−3/2=−12(σ2+ϵ)∑ndx^nx^nd\sigma^2 = \sum_n d\hat x_n\cdot\big(-\tfrac12\big)(x_n-\mu)(\sigma^2+\epsilon)^{-3/2} = -\dfrac{1}{2(\sigma^2+\epsilon)}\sum_n d\hat x_n\hat x_n.The chain rule sums over the NN children of σ2\sigma^2, and (xn−μ)(σ2+ϵ)−1/2=x^n(x_n-\mu)(\sigma^2+\epsilon)^{-1/2} = \hat x_n leaves one factor (σ2+ϵ)−1(\sigma^2+\epsilon)^{-1}.
  3. μ\mu feeds every x^n\hat x_n, with ∂x^n/∂μ=−1/σ2+ϵ\partial\hat x_n/\partial\mu = -1/\sqrt{\sigma^2+\epsilon}, and feeds σ2\sigma^2, with ∂σ2/∂μ=−2N∑n(xn−μ)\partial\sigma^2/\partial\mu = -\tfrac2N\sum_n (x_n-\mu).Differentiate each child of μ\mu with its other inputs held fixed.
  4. ∂σ2/∂μ=0\partial\sigma^2/\partial\mu = 0.Deviations from the mean sum to zero (Problem 3, step 4).
  5. dμ=−1σ2+ϵ∑ndx^n+dσ2⋅0d\mu = -\dfrac{1}{\sqrt{\sigma^2+\epsilon}}\sum_n d\hat x_n + d\sigma^2\cdot 0.The chain rule sums over both kinds of child, the x^n\hat x_n and σ2\sigma^2.
  6. dσ2=−12(σ2+ϵ)∑ndx^nx^nd\sigma^2 = -\dfrac{1}{2(\sigma^2+\epsilon)}\textstyle\sum_n d\hat x_n\hat x_n; dμ=−1σ2+ϵ∑ndx^nd\mu = -\dfrac{1}{\sqrt{\sigma^2+\epsilon}}\textstyle\sum_n d\hat x_nTwo scalars per feature. The path from μ\mu through σ2\sigma^2 is real but carries nothing, for the same reason that ∇xσ2\nabla_x\sigma^2 in Problem 3 has no μ\mu term: μ\mu minimises the mean squared deviation, so a small change in it moves σ2\sigma^2 only to second order.

Problem 5

Assemble dxdx for one feature from dx^d\hat x, dσ2d\sigma^2 and dμd\mu, and simplify it to a form that uses only dx^d\hat x, x^\hat x and σ2\sigma^2.

  1. xnx_n feeds x^n\hat x_n with ∂x^n/∂xn=1/σ2+ϵ\partial\hat x_n/\partial x_n = 1/\sqrt{\sigma^2+\epsilon}, feeds σ2\sigma^2 with ∂σ2/∂xn=2N(xn−μ)\partial\sigma^2/\partial x_n = \tfrac2N(x_n-\mu), and feeds μ\mu with ∂μ/∂xn=1N\partial\mu/\partial x_n = \tfrac1N.These are the children of xnx_n in the graph of Problem 4, each differentiated with its other inputs held fixed; xnx_n reaches x^m\hat x_m for m≠nm \ne n only through μ\mu and σ2\sigma^2.
  2. dxn=dx^nσ2+ϵ+dσ2⋅2N(xn−μ)+dμNdx_n = \dfrac{d\hat x_n}{\sqrt{\sigma^2+\epsilon}} + d\sigma^2\cdot\dfrac2N(x_n-\mu) + \dfrac{d\mu}{N}.The chain rule sums over the three children.
  3. dσ2⋅2N(xn−μ)=−1N(σ2+ϵ)(xn−μ)∑mdx^mx^m=−1Nσ2+ϵ x^n∑mdx^mx^md\sigma^2\cdot\dfrac2N(x_n-\mu) = -\dfrac{1}{N(\sigma^2+\epsilon)}(x_n-\mu)\sum_m d\hat x_m\hat x_m = -\dfrac{1}{N\sqrt{\sigma^2+\epsilon}}\,\hat x_n\sum_m d\hat x_m\hat x_m.Problem 4 for dσ2d\sigma^2, with the summation index renamed mm; then (xn−μ)/σ2+ϵ=x^n(x_n-\mu)/\sqrt{\sigma^2+\epsilon} = \hat x_n.
  4. dμN=−1Nσ2+ϵ∑mdx^m\dfrac{d\mu}{N} = -\dfrac{1}{N\sqrt{\sigma^2+\epsilon}}\sum_m d\hat x_m.Problem 4 for dμd\mu.
  5. dxn=1Nσ2+ϵ(N dx^n−∑mdx^m−x^n∑mdx^mx^m)dx_n = \dfrac{1}{N\sqrt{\sigma^2+\epsilon}}\Big(N\,d\hat x_n - \sum_m d\hat x_m - \hat x_n\sum_m d\hat x_m\hat x_m\Big).Put the three terms over the common factor 1/(Nσ2+ϵ)1/(N\sqrt{\sigma^2+\epsilon}); the first term becomes N dx^nN\,d\hat x_n.
  6. dx=1Nσ2+ϵ(N dx^−(∑ndx^n)1−x^∑ndx^nx^n)dx = \dfrac{1}{N\sqrt{\sigma^2+\epsilon}}\Big(N\,d\hat x - \big(\textstyle\sum_n d\hat x_n\big)\mathbf{1} - \hat x\textstyle\sum_n d\hat x_n\hat x_n\Big) for each featureN×1N\times 1. For all features at once, dX=1N(N dX^−11⊤dX^−X^⊙11⊤(dX^⊙X^))DdX = \tfrac1N\big(N\,d\hat X - \mathbf{1}\mathbf{1}^\top d\hat X - \hat X\odot\mathbf{1}\mathbf{1}^\top(d\hat X\odot\hat X)\big)D, where 11⊤M\mathbf{1}\mathbf{1}^\top M puts the column sums of MM in every row. Two sums per feature and elementwise work: O(Nd)O(Nd), with X^\hat X and σ2\sigma^2 saved from the forward pass.

Problem 6

For one feature, show that 1⊤dx=0\mathbf{1}^\top dx = 0 and compute x^⊤dx\hat x^\top dx. When is dxdx orthogonal to x^\hat x?

  1. 1⊤x^=∑n(xn−μ)/σ2+ϵ=0\mathbf{1}^\top\hat x = \sum_n (x_n-\mu)/\sqrt{\sigma^2+\epsilon} = 0.Deviations from the mean sum to zero, whatever ϵ\epsilon is.
  2. x^⊤x^=∑n(xn−μ)2σ2+ϵ=Nσ2σ2+ϵ\hat x^\top\hat x = \dfrac{\sum_n (x_n-\mu)^2}{\sigma^2+\epsilon} = \dfrac{N\sigma^2}{\sigma^2+\epsilon}.∑n(xn−μ)2=Nσ2\sum_n (x_n-\mu)^2 = N\sigma^2 by the definition of the biased variance.
  3. 1⊤dx=1Nσ2+ϵ(N∑ndx^n−N∑ndx^n−(1⊤x^)∑ndx^nx^n)=0\mathbf{1}^\top dx = \dfrac{1}{N\sqrt{\sigma^2+\epsilon}}\Big(N\sum_n d\hat x_n - N\sum_n d\hat x_n - (\mathbf{1}^\top\hat x)\sum_n d\hat x_n\hat x_n\Big) = 0.Problem 5 with 1⊤1=N\mathbf{1}^\top\mathbf{1} = N, and step 1 for the last term.
  4. x^⊤dx=1Nσ2+ϵ(N−Nσ2σ2+ϵ)∑ndx^nx^n\hat x^\top dx = \dfrac{1}{N\sqrt{\sigma^2+\epsilon}}\Big(N - \dfrac{N\sigma^2}{\sigma^2+\epsilon}\Big)\sum_n d\hat x_n\hat x_n.Problem 5 again: x^⊤dx^=∑ndx^nx^n\hat x^\top d\hat x = \sum_n d\hat x_n\hat x_n, the middle term vanishes by step 1, and step 2 gives the last.
  5. 1⊤dx=0\mathbf{1}^\top dx = 0; x^⊤dx=1σ2+ϵ ϵσ2+ϵ∑ndx^nx^n\hat x^\top dx = \dfrac{1}{\sqrt{\sigma^2+\epsilon}}\,\dfrac{\epsilon}{\sigma^2+\epsilon}\textstyle\sum_n d\hat x_n\hat x_n, which is 00 when ϵ=0\epsilon = 0N−Nσ2/(σ2+ϵ)=Nϵ/(σ2+ϵ)N - N\sigma^2/(\sigma^2+\epsilon) = N\epsilon/(\sigma^2+\epsilon). Adding a constant to one feature across the batch leaves X^\hat X unchanged, and so (up to ϵ\epsilon) does scaling it, so the loss cannot push either way. It is the layer-norm page's Problem 8 with columns in place of rows. One consequence: a bias added to a feature just before batch norm gets gradient 1⊤dx=0\mathbf{1}^\top dx = 0 and never moves, which is why layers followed by batch norm usually drop their bias.

Problem 7

Let LN⁡(Z)\operatorname{LN}(Z) be layer norm without γ\gamma and β\beta, applied to each row of a matrix ZZ with that row's own mean, biased variance and ϵ\epsilon. Express X^\hat X and dXdX through LN⁡\operatorname{LN}, and say which axis each layer averages over.

  1. Row jj of X⊤X^\top is column jj of XX: the NN values of feature jj.Transposing swaps the roles of rows and columns.
  2. Row jj of LN⁡(X⊤)\operatorname{LN}(X^\top) is ((x−μj1)/σj2+ϵ)⊤\big((x - \mu_j\mathbf{1})/\sqrt{\sigma_j^2+\epsilon}\big)^\top, which is column jj of X^\hat X as a row.Layer norm's row mean and biased variance, taken over the NN entries of that row, are μj\mu_j and σj2\sigma_j^2.
  3. So X^=LN⁡(X⊤)⊤\hat X = \operatorname{LN}(X^\top)^\top, and LL depends on XX through X⊤X^\top, then LN⁡\operatorname{LN}, then a transpose.Step 2 for every jj.
  4. A transpose only relabels entries, so its backward pass is a transpose: dXdX is the transpose of the gradient at X⊤X^\top, and the upstream gradient at LN⁡(X⊤)\operatorname{LN}(X^\top) is dX^⊤d\hat X^\top.Each entry of XX is one entry of X⊤X^\top with derivative 11, and likewise for X^\hat X.
  5. The layer-norm page's Problem 7 on row jj of X⊤X^\top, with upstream dx^⊤d\hat x^\top and γ=1\gamma = \mathbf{1}, gives 1σj2+ϵ(dx^−1N(∑ndx^n)1−x^ 1N∑ndx^nx^n)\tfrac{1}{\sqrt{\sigma_j^2+\epsilon}}\big(d\hat x - \tfrac1N(\sum_n d\hat x_n)\mathbf{1} - \hat x\,\tfrac1N\sum_n d\hat x_n\hat x_n\big), as a row.Its mean⁡\operatorname{mean} is 1N∑n\tfrac1N\sum_n here; factoring out 1N\tfrac1N gives Problem 5 exactly.
  6. X^=LN⁡(X⊤)⊤\hat X = \operatorname{LN}(X^\top)^\top, and dXdX is the transpose of layer norm's input gradient at X⊤X^\top with upstream dX^⊤d\hat X^\top: batch norm takes its means down each column, over the NN examples (axis 0), and layer norm along each row, over the dd features (axis 1)The algebra is the same; what changes is which entries share statistics. In layer norm row nn of dXdX depends only on row nn of dX^d\hat X, so examples stay independent. In batch norm column jj of dXdX depends on all of column jj of dX^d\hat X, so each example's gradient depends on the rest of the batch.

Problem 8

At inference the layer uses the running statistics: X^=(X−1μˉ⊤)Dˉ\hat X = (X - \mathbf{1}\bar\mu^\top)\bar D. Compute dXdX, dγd\gamma and dβd\beta.

  1. X^nj=(Xnj−μˉj)/σˉj2+ϵ\hat X_{nj} = (X_{nj} - \bar\mu_j)/\sqrt{\bar\sigma_j^2+\epsilon}, so X^nj\hat X_{nj} depends only on XnjX_{nj}, with slope (σˉj2+ϵ)−1/2(\bar\sigma_j^2+\epsilon)^{-1/2}.μˉ\bar\mu and σˉ2\bar\sigma^2 were accumulated from earlier batches and do not depend on the current XX.
  2. Ynj=γjX^nj+βjY_{nj} = \gamma_j\hat X_{nj} + \beta_j depends only on XnjX_{nj}, with slope γj(σˉj2+ϵ)−1/2\gamma_j(\bar\sigma_j^2+\epsilon)^{-1/2}.Step 1 and the chain rule through one scalar.
  3. ∂L/∂Xnj=γj dYnj/σˉj2+ϵ\partial L/\partial X_{nj} = \gamma_j\,dY_{nj}/\sqrt{\bar\sigma_j^2+\epsilon}, and over all entries that is dYdiag⁡(γ)DˉdY\operatorname{diag}(\gamma)\bar D.XnjX_{nj} reaches LL only through YnjY_{nj}; right-multiplying by the diagonal matrices scales column jj by γj\gamma_j and by (σˉj2+ϵ)−1/2(\bar\sigma_j^2+\epsilon)^{-1/2}.
  4. dγd\gamma and dβd\beta come from Problem 2 unchanged, with this X^\hat X.Problem 2 used only Y=X^diag⁡(γ)+1β⊤Y = \hat X\operatorname{diag}(\gamma) + \mathbf{1}\beta^\top, not where X^\hat X came from.
  5. dX=dYdiag⁡(γ)DˉdX = dY\operatorname{diag}(\gamma)\bar D, that is dXnj=γj dYnj/σˉj2+ϵdX_{nj} = \gamma_j\,dY_{nj}/\sqrt{\bar\sigma_j^2+\epsilon}; dγ=(dY⊙X^)⊤1d\gamma = (dY\odot\hat X)^\top\mathbf{1} and dβ=dY⊤1d\beta = dY^\top\mathbf{1} with X^=(X−1μˉ⊤)Dˉ\hat X = (X - \mathbf{1}\bar\mu^\top)\bar DN×dN\times d, dd and dd. With fixed statistics the layer is an elementwise affine map, Y=XDˉdiag⁡(γ)+1(β−diag⁡(γ)Dˉμˉ)⊤Y = X\bar D\operatorname{diag}(\gamma) + \mathbf{1}\big(\beta - \operatorname{diag}(\gamma)\bar D\bar\mu\big)^\top, so examples no longer interact, its gradient has no centring terms, and it can be folded into the weights and bias of the layer before it. This is the gradient used when a network is fine-tuned with batch norm frozen.

Problem 9

Train with a batch of one (N=1N = 1). Compute μ\mu, σ2\sigma^2, X^\hat X, YY and the gradients dXdX, dγd\gamma, dβd\beta.

  1. μ=X⊤\mu = X^\top, the single row as a column, and σ2=0\sigma^2 = 0.Each column has one entry, which is its own mean and has zero deviation from it.
  2. X^=(X−1μ⊤)D=0⋅D=0\hat X = (X - \mathbf{1}\mu^\top)D = 0\cdot D = 0, and D=ϵ−1/2ID = \epsilon^{-1/2}I is finite.X−1μ⊤=X−X=0X - \mathbf{1}\mu^\top = X - X = 0; with ϵ=0\epsilon = 0 every entry would be 0/00/0.
  3. Y=0⋅diag⁡(γ)+1β⊤=β⊤Y = 0\cdot\operatorname{diag}(\gamma) + \mathbf{1}\beta^\top = \beta^\top.1\mathbf{1} has one entry.
  4. YY does not depend on XX or γ\gamma at all, so dX=0dX = 0 and dγ=(dY⊙0)⊤1=0d\gamma = (dY\odot 0)^\top\mathbf{1} = 0, while dβ=dY⊤1=dY⊤d\beta = dY^\top\mathbf{1} = dY^\top.Problem 2. Problem 5 agrees: with N=1N = 1 and x^=0\hat x = 0 it gives 1ϵ(dx^−dx^−0)=0\tfrac{1}{\sqrt\epsilon}(d\hat x - d\hat x - 0) = 0.
  5. With N=1N = 1: μ=X⊤\mu = X^\top, σ2=0\sigma^2 = 0, X^=0\hat X = 0, Y=β⊤Y = \beta^\top; dX=0dX = 0, dγ=0d\gamma = 0, dβ=dY⊤d\beta = dY^\topThe layer outputs β\beta whatever the input, so nothing below it receives a gradient and γ\gamma never trains. Small batches degrade the same way more gently: the statistics of a few examples are noisy, and every example's gradient is entangled with them.

Problem 10

With N≥2N \ge 2, one feature takes the same value on every example in the batch. Compute x^\hat x and dxdx for that feature. What does ϵ\epsilon do here, and how large is the effect?

  1. μ\mu equals that common value, so x−μ1=0x - \mu\mathbf{1} = 0 and σ2=0\sigma^2 = 0.The mean of equal numbers is that number, and every deviation is zero.
  2. x^=0/0+ϵ=0\hat x = 0/\sqrt{0+\epsilon} = 0.ϵ>0\epsilon > 0 makes the denominator ϵ\sqrt\epsilon; with ϵ=0\epsilon = 0 it would be 0/00/0.
  3. Problem 5 holds at σ2=0\sigma^2 = 0: dx=1Nϵ(N dx^−(∑ndx^n)1−0)dx = \tfrac{1}{N\sqrt\epsilon}\big(N\,d\hat x - (\sum_n d\hat x_n)\mathbf{1} - 0\big).Its derivation needed only σ2+ϵ>0\sigma^2+\epsilon > 0, so the map is smooth here; the last term carries the factor x^=0\hat x = 0.
  4. On a constant feature x^=0\hat x = 0 and dx=1ϵ(dx^−1N(∑ndx^n)1)dx = \dfrac{1}{\sqrt\epsilon}\Big(d\hat x - \dfrac1N\big(\textstyle\sum_n d\hat x_n\big)\mathbf{1}\Big); without ϵ\epsilon the forward pass would divide 00 by 00Distribute the 1N\tfrac1N over the bracket. The gain is 1/ϵ1/\sqrt\epsilon: about 316316 for ϵ=10−5\epsilon = 10^{-5}, so the centred upstream gradient comes back several hundred times larger, and a tiny spread in a nearly constant feature is blown up to unit scale in the forward pass. For a feature with σ2≫ϵ\sigma^2 \gg \epsilon, ϵ\epsilon changes x^\hat x only by a relative ϵ/(2σ2)\epsilon/(2\sigma^2), and the x^\hat x component of dxdx in Problem 6 is of the same small order.

Where this goes wrong

1. Taking the mean over the features instead of the batch

Array code computes a mean with an axis argument, and layer norm, the normalisation layer of every transformer, takes it along the last axis.

  1. X∈RN×dX \in \mathbb{R}^{N\times d} with rows as examplesRight so far: the layout of Problem 1.
  2. “Normalise using the mean over the last axis, as layer norm does.”The habit that causes the mistake: layer norm's axis carried over, when for rows-as-examples the last axis is the features.
  3. μ=1dX1∈RN\mu = \tfrac1d X\mathbf{1} \in \mathbb{R}^N, one mean per exampleThat is layer norm (Problem 7). Batch norm needs μ=1NX⊤1∈Rd\mu = \tfrac1N X^\top\mathbf{1} \in \mathbb{R}^d, one per feature. When N=dN = d the wrong μ\mu has the right length, broadcasts against the rows without error, and subtracts example jj's mean from feature jj.

2. Holding the batch mean and variance constant

At inference the gradient is elementwise (Problem 8), and the training forward pass looks the same with batch statistics in place of running ones.

  1. dX^=dYdiag⁡(γ)d\hat X = dY\operatorname{diag}(\gamma)Right so far: Problem 2.
  2. “Training differs from inference only in which μ\mu and σ2\sigma^2 are used, so the backward pass is Problem 8 with DD in place of Dˉ\bar D.”The analogy that causes the mistake: running statistics are constants, but the batch statistics are computed from the current XX, so they are functions of it.
  3. dX=dYdiag⁡(γ)DdX = dY\operatorname{diag}(\gamma)DIt keeps only the first term of Problem 5 and drops the paths through μ\mu and σ2\sigma^2. Its column sums are (1⊤dX^)D\big(\mathbf{1}^\top d\hat X\big)D, not 00 (Problem 6), so it claims that shifting a feature across the batch changes the loss, which it cannot.

3. Differentiating the unbiased variance against a biased forward pass

Statistics courses, and the default variance function of some array libraries, divide by N−1N - 1, and a hand-written backward pass often re-derives ∂σ2/∂x\partial\sigma^2/\partial x from that formula.

  1. dxn=dx^nσ2+ϵ+dσ2 ∂σ2∂xn+dμNdx_n = \dfrac{d\hat x_n}{\sqrt{\sigma^2+\epsilon}} + d\sigma^2\,\dfrac{\partial\sigma^2}{\partial x_n} + \dfrac{d\mu}{N}Right so far: Problem 5, step 2, with dσ2d\sigma^2 and dμd\mu from Problem 4.
  2. “The sample variance is 1N−1∑n(xn−μ)2\tfrac{1}{N-1}\sum_n (x_n-\mu)^2, so ∂σ2/∂xn=2N−1(xn−μ)\partial\sigma^2/\partial x_n = \tfrac{2}{N-1}(x_n-\mu).”The shortcut that causes the mistake: the derivative of the unbiased estimator, while the forward pass divided by NN (Problem 3).
  3. dx=1Nσ2+ϵ(N dx^−(∑ndx^n)1−NN−1 x^∑ndx^nx^n)dx = \dfrac{1}{N\sqrt{\sigma^2+\epsilon}}\Big(N\,d\hat x - \big(\textstyle\sum_n d\hat x_n\big)\mathbf{1} - \dfrac{N}{N-1}\,\hat x\textstyle\sum_n d\hat x_n\hat x_n\Big)It is the derivative of a different function from the one the forward pass computed: the x^\hat x term is N/(N−1)N/(N-1) times too large. It still sums to zero, so a sum-to-zero test passes, and at batch size 256256 the error is under 0.4%0.4\%, small enough to slip through a loose gradient check.

4. Leaving γ's gradient unsummed over the batch

Each example contributes its own term to γ\gamma's gradient, and the array of those terms already has a familiar shape.

  1. ∂L/∂γj=∑ndYnjX^nj\partial L/\partial\gamma_j = \sum_n dY_{nj}\hat X_{nj}Right so far: Problem 2, step 2.
  2. “Each example's contribution is dY⊙X^dY\odot\hat X in its own row, so the gradient is that array.”The shortcut that causes the mistake: stopping at the per-example contributions, the way dX^d\hat X is per example, when γ\gamma is a parameter shared by all of them.
  3. dγ=dY⊙X^d\gamma = dY\odot\hat XIt is N×dN\times d, not dd: the sum over examples is missing, and the gradient is (dY⊙X^)⊤1(dY\odot\hat X)^\top\mathbf{1} (Problem 2). In array code the update γ−η dY⊙X^\gamma - \eta\,dY\odot\hat X, with learning rate η\eta, broadcasts silently and turns γ\gamma into one scale per example.

5. Scaling the training gradient by the running variance

A batch-norm layer keeps μˉ\bar\mu and σˉ2\bar\sigma^2 as stored attributes, while the batch statistics of the forward pass are temporaries that hand-written code has to remember to save.

  1. In training, X^=(X−1μ⊤)D\hat X = (X - \mathbf{1}\mu^\top)D with the batch statistics μ\mu and σ2\sigma^2Right so far: the training forward pass as defined before Problem 1.
  2. “The layer's variance is σˉ2\bar\sigma^2, so use it for the 1/σ2+ϵ1/\sqrt{\sigma^2+\epsilon} in the backward pass.”The shortcut that causes the mistake: reading the stored statistic instead of saving the one the forward pass divided by.
  3. dx=1Nσˉ2+ϵ(N dx^−(∑ndx^n)1−x^∑ndx^nx^n)dx = \dfrac{1}{N\sqrt{\bar\sigma^2+\epsilon}}\Big(N\,d\hat x - \big(\textstyle\sum_n d\hat x_n\big)\mathbf{1} - \hat x\textstyle\sum_n d\hat x_n\hat x_n\Big)The factor comes from differentiating the training forward pass, which divided by σ2+ϵ\sqrt{\sigma^2+\epsilon} with the batch variance (Problem 5); σˉ2\bar\sigma^2 is not on the path from XX to LL in training at all. The answer is off by the factor (σ2+ϵ)/(σˉ2+ϵ)\sqrt{(\sigma^2+\epsilon)/(\bar\sigma^2+\epsilon)}, close to 11 once the running average has settled, so it hides, and far from 11 early in training, when σˉ2\bar\sigma^2 still holds its initial value.

Print this set: batch-norm-backward.pdf (problems, answers, and worked solutions on separate pages).