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 γ and .β. 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, γ'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, ⊙ is the elementwise product, 1 is the all-ones vector, and rows are examples. As on the previous page, dA is the code-style name for ,∇AL, with the shape of .A.
The input is :X∈RN×d:N examples, d features. Batch norm works down the columns. Feature j has batch mean μj=N1∑nXnj and biased variance ,σj2=N1∑n(Xnj−μj)2, and μ,σ2∈Rd collect them.
ϵ>0 is a small constant inside the square root. D is the d×d diagonal matrix with ,Djj=(σj2+ϵ)−1/2, so the normalised input is ,X^=(X−1μ⊤)D, with entries .X^nj=(Xnj−μj)/σj2+ϵ.
The output is ,Y=X^diag(γ)+1β⊤, that is ,Ynj=γjX^nj+βj, with learned .γ,β∈Rd.L is a scalar loss that depends on ,X,γ and β only through ,Y, and dY is given.
One feature at a time: x∈RN is column j of ,X,μ and σ2 are then the scalars μj and ,σj2,x^=(x−μ1)/σ2+ϵ is column j of ,X^, and ,dx^,dx are column j of ,dX^,.dX. Sums ∑n run over the N examples.
Running statistics μˉ,σˉ2∈Rd are moving averages of the batch statistics, updated outside the gradient; at inference they replace μ and ,σ2, and Dˉ is D built from .σˉ2.
Give the shapes of ,μ,,σ2,,X^,,γ,β and .Y. Which entries of X does Ynj depend on?
·
Compute ,dγ,dβ and dX^ from .dY.
··
For one feature ,x∈RN, compute ∇xμ and .∇xσ2.
··
Treat one feature's forward pass as a graph: μ is computed from ,x,σ2=N1∑n(xn−μ)2 from x and ,μ, and x^n=(xn−μ)/σ2+ϵ from ,x,μ and .σ2. Given ,dx^, compute the gradients dσ2=∂L/∂σ2 and dμ=∂L/∂μ arriving at those two nodes.
···
Assemble dx for one feature from ,dx^,dσ2 and ,dμ, and simplify it to a form that uses only ,dx^,x^ and .σ2.
··
For one feature, show that 1⊤dx=0 and compute .x^⊤dx. When is dx orthogonal to ?x^?
··
Let LN(Z) be layer norm without γ and ,β, applied to each row of a matrix Z with that row's own mean, biased variance and .ϵ. Express X^ and dX through ,LN, and say which axis each layer averages over.
··
At inference the layer uses the running statistics: .X^=(X−1μˉ⊤)Dˉ. Compute ,dX,dγ and .dβ.
·
Train with a batch of one ().N=1). Compute ,μ,,σ2,,X^,Y and the gradients ,dX,,dγ,.dβ.
···
With ,N≥2, one feature takes the same value on every example in the batch. Compute x^ and dx for that feature. What does ϵ do here, and how large is the effect?
Answers
;μ,σ2,γ,β∈Rd;;X^,Y∈RN×d;Ynj depends on every entry of column j of X and on no other column
;dγ=(dY⊙X^)⊤1;;dβ=dY⊤1;dX^=dYdiag(γ)
;∇xμ=N11;∇xσ2=N2(x−μ1)
;dσ2=−2(σ2+ϵ)1∑ndx^nx^n;dμ=−σ2+ϵ1∑ndx^n
dx=Nσ2+ϵ1(Ndx^−(∑ndx^n)1−x^∑ndx^nx^n) for each feature
;1⊤dx=0;,x^⊤dx=σ2+ϵ1σ2+ϵϵ∑ndx^nx^n, which is 0 when ϵ=0
,X^=LN(X⊤)⊤, and dX is the transpose of layer norm's input gradient at X⊤ with upstream :dX^⊤: batch norm takes its means down each column, over the N examples (axis 0), and layer norm along each row, over the d features (axis 1)
,dX=dYdiag(γ)Dˉ, that is ;dXnj=γjdYnj/σˉj2+ϵ;dγ=(dY⊙X^)⊤1 and dβ=dY⊤1 with X^=(X−1μˉ⊤)Dˉ
With :N=1:,μ=X⊤,,σ2=0,,X^=0,;Y=β⊤;,dX=0,,dγ=0,dβ=dY⊤
On a constant feature x^=0 and ;dx=ϵ1(dx^−N1(∑ndx^n)1); without ϵ the forward pass would divide 0 by 0
Worked solutions
Problem 1
Give the shapes of ,μ,,σ2,,X^,,γ,β and .Y. Which entries of X does Ynj depend on?
,μ=N1X⊤1∈Rd, and likewise .σ2∈Rd.Each is a sum over the N examples, which leaves one number per feature; .(d×N)(N×1)=d×1.
X−1μ⊤ is N×d and subtracts μj from every entry of column ;j; multiplying on the right by the diagonal D scales column j by ,(σj2+ϵ)−1/2, so X^ is .N×d.,(1μ⊤)nj=μj, and (AD)nj=AnjDjj for a diagonal .D.
Y=X^diag(γ)+1β⊤ is ,N×d, with .γ,β∈Rd.One scale and one shift per feature, shared by every example, as the bias of the previous page's dense layer is.
X^nj uses ,Xnj,μj and ,σj2, and μj and σj2 are built from all N entries of column j and nothing else.The definitions of μj and σj2 sum over n with j fixed.
;μ,σ2,γ,β∈Rd;;X^,Y∈RN×d;Ynj depends on every entry of column j of X and on no other columnFeatures never mix and examples always do, the reverse of the dense layer on the previous page, where row n of the output used only row n 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β and dX^ from .dY.
.Ynj=γjX^nj+βj.Index form of ;Y=X^diag(γ)+1β⊤;X^ does not depend on γ or .β.
∂L/∂γj=∑ndYnjX^nj and .∂L/∂βj=∑ndYnj.γj and βj appear in Ynj for every example ,n, with coefficients X^nj and ,1, and the chain rule sums over every entry that contains them.
.∂L/∂X^nj=γjdYnj.X^nj appears only in ,Ynj, with coefficient .γj.
For a matrix M with N rows, M⊤1 adds up its rows, and Mdiag(γ) scales column j by .γj.(M⊤1)j=∑nMnj and .(Mdiag(γ))nj=Mnjγj.
;dγ=(dY⊙X^)⊤1;;dβ=dY⊤1;dX^=dYdiag(γ)Shapes ,d,d and .N×d. The sums over the batch appear because γ and β are shared by every example. From here on the input gradient needs only ,dX^, one column at a time: for feature ,j,dx^=γj times column j of .dY.
Problem 3
For one feature ,x∈RN, compute ∇xμ and .∇xσ2.
∂μ/∂xm=N1 for every .m.μ=N1∑nxn contains xm once, with coefficient .N1.
.∂σ2/∂xm=N2∑n(xn−μ)(∂xn/∂xm−∂μ/∂xm).Chain rule on each square ,(xn−μ)2, with μ kept as the function of x that it is.
.∂σ2/∂xm=N2(xm−μ)−N22∑n(xn−μ).∂xn/∂xm is 1 for n=m and 0 otherwise, which picks out one term; step 1 gives ,∂μ/∂xm=N1, the same for every .n.
.∑n(xn−μ)=∑nxn−Nμ=0.Nμ=∑nxn by the definition of .μ.
;∇xμ=N11;∇xσ2=N2(x−μ1)Both ,N×1, the shape of .x. The N1 in ∇xσ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: μ is computed from ,x,σ2=N1∑n(xn−μ)2 from x and ,μ, and x^n=(xn−μ)/σ2+ϵ from ,x,μ and .σ2. Given ,dx^, compute the gradients dσ2=∂L/∂σ2 and dμ=∂L/∂μ arriving at those two nodes.
σ2 feeds only the ,x^n, and .∂x^n/∂σ2=−21(xn−μ)(σ2+ϵ)−3/2.In the graph each node's inputs are held fixed when it is varied; power rule on .(σ2+ϵ)−1/2.
.dσ2=∑ndx^n⋅(−21)(xn−μ)(σ2+ϵ)−3/2=−2(σ2+ϵ)1∑ndx^nx^n.The chain rule sums over the N children of ,σ2, and (xn−μ)(σ2+ϵ)−1/2=x^n leaves one factor .(σ2+ϵ)−1.
μ feeds every ,x^n, with ,∂x^n/∂μ=−1/σ2+ϵ, and feeds ,σ2, with .∂σ2/∂μ=−N2∑n(xn−μ).Differentiate each child of μ with its other inputs held fixed.
.∂σ2/∂μ=0.Deviations from the mean sum to zero (Problem 3, step 4).
.dμ=−σ2+ϵ1∑ndx^n+dσ2⋅0.The chain rule sums over both kinds of child, the x^n and .σ2.
;dσ2=−2(σ2+ϵ)1∑ndx^nx^n;dμ=−σ2+ϵ1∑ndx^nTwo scalars per feature. The path from μ through σ2 is real but carries nothing, for the same reason that ∇xσ2 in Problem 3 has no μ term: μ minimises the mean squared deviation, so a small change in it moves σ2 only to second order.
Problem 5
Assemble dx for one feature from ,dx^,dσ2 and ,dμ, and simplify it to a form that uses only ,dx^,x^ and .σ2.
xn feeds x^n with ,∂x^n/∂xn=1/σ2+ϵ, feeds σ2 with ,∂σ2/∂xn=N2(xn−μ), and feeds μ with .∂μ/∂xn=N1.These are the children of xn in the graph of Problem 4, each differentiated with its other inputs held fixed; xn reaches x^m for m=n only through μ and .σ2.
.dxn=σ2+ϵdx^n+dσ2⋅N2(xn−μ)+Ndμ.The chain rule sums over the three children.
.dσ2⋅N2(xn−μ)=−N(σ2+ϵ)1(xn−μ)∑mdx^mx^m=−Nσ2+ϵ1x^n∑mdx^mx^m.Problem 4 for ,dσ2, with the summation index renamed ;m; then .(xn−μ)/σ2+ϵ=x^n.
.Ndμ=−Nσ2+ϵ1∑mdx^m.Problem 4 for .dμ.
.dxn=Nσ2+ϵ1(Ndx^n−∑mdx^m−x^n∑mdx^mx^m).Put the three terms over the common factor ;1/(Nσ2+ϵ); the first term becomes .Ndx^n.
dx=Nσ2+ϵ1(Ndx^−(∑ndx^n)1−x^∑ndx^nx^n) for each feature.N×1. For all features at once, ,dX=N1(NdX^−11⊤dX^−X^⊙11⊤(dX^⊙X^))D, where 11⊤M puts the column sums of M in every row. Two sums per feature and elementwise work: ,O(Nd), with X^ and σ2 saved from the forward pass.
Problem 6
For one feature, show that 1⊤dx=0 and compute .x^⊤dx. When is dx orthogonal to ?x^?
.1⊤x^=∑n(xn−μ)/σ2+ϵ=0.Deviations from the mean sum to zero, whatever ϵ is.
.x^⊤x^=σ2+ϵ∑n(xn−μ)2=σ2+ϵNσ2.∑n(xn−μ)2=Nσ2 by the definition of the biased variance.
.1⊤dx=Nσ2+ϵ1(N∑ndx^n−N∑ndx^n−(1⊤x^)∑ndx^nx^n)=0.Problem 5 with ,1⊤1=N, and step 1 for the last term.
.x^⊤dx=Nσ2+ϵ1(N−σ2+ϵNσ2)∑ndx^nx^n.Problem 5 again: ,x^⊤dx^=∑ndx^nx^n, the middle term vanishes by step 1, and step 2 gives the last.
;1⊤dx=0;,x^⊤dx=σ2+ϵ1σ2+ϵϵ∑ndx^nx^n, which is 0 when ϵ=0.N−Nσ2/(σ2+ϵ)=Nϵ/(σ2+ϵ). Adding a constant to one feature across the batch leaves X^ unchanged, and so (up to )ϵ) 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 and never moves, which is why layers followed by batch norm usually drop their bias.
Problem 7
Let LN(Z) be layer norm without γ and ,β, applied to each row of a matrix Z with that row's own mean, biased variance and .ϵ. Express X^ and dX through ,LN, and say which axis each layer averages over.
Row j of X⊤ is column j of :X: the N values of feature .j.Transposing swaps the roles of rows and columns.
Row j of LN(X⊤) is ,((x−μj1)/σj2+ϵ)⊤, which is column j of X^ as a row.Layer norm's row mean and biased variance, taken over the N entries of that row, are μj and .σj2.
So ,X^=LN(X⊤)⊤, and L depends on X through ,X⊤, then ,LN, then a transpose.Step 2 for every .j.
A transpose only relabels entries, so its backward pass is a transpose: dX is the transpose of the gradient at ,X⊤, and the upstream gradient at LN(X⊤) is .dX^⊤.Each entry of X is one entry of X⊤ with derivative ,1, and likewise for .X^.
The layer-norm page's Problem 7 on row j of ,X⊤, with upstream dx^⊤ and ,γ=1, gives ,σj2+ϵ1(dx^−N1(∑ndx^n)1−x^N1∑ndx^nx^n), as a row.Its mean is N1∑n here; factoring out N1 gives Problem 5 exactly.
,X^=LN(X⊤)⊤, and dX is the transpose of layer norm's input gradient at X⊤ with upstream :dX^⊤: batch norm takes its means down each column, over the N examples (axis 0), and layer norm along each row, over the d features (axis 1)The algebra is the same; what changes is which entries share statistics. In layer norm row n of dX depends only on row n of ,dX^, so examples stay independent. In batch norm column j of dX depends on all of column j of ,dX^, 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ˉ. Compute ,dX,dγ and .dβ.
,X^nj=(Xnj−μˉj)/σˉj2+ϵ, so X^nj depends only on ,Xnj, with slope .(σˉj2+ϵ)−1/2.μˉ and σˉ2 were accumulated from earlier batches and do not depend on the current .X.
Ynj=γjX^nj+βj depends only on ,Xnj, with slope .γj(σˉj2+ϵ)−1/2.Step 1 and the chain rule through one scalar.
,∂L/∂Xnj=γjdYnj/σˉj2+ϵ, and over all entries that is .dYdiag(γ)Dˉ.Xnj reaches L only through ;Ynj; right-multiplying by the diagonal matrices scales column j by γj and by .(σˉj2+ϵ)−1/2.
dγ and dβ come from Problem 2 unchanged, with this .X^.Problem 2 used only ,Y=X^diag(γ)+1β⊤, not where X^ came from.
,dX=dYdiag(γ)Dˉ, that is ;dXnj=γjdYnj/σˉj2+ϵ;dγ=(dY⊙X^)⊤1 and dβ=dY⊤1 with X^=(X−1μˉ⊤)Dˉ,N×d,d and .d. With fixed statistics the layer is an elementwise affine map, ,Y=XDˉdiag(γ)+1(β−diag(γ)Dˉμˉ)⊤, 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=1). Compute ,μ,,σ2,,X^,Y and the gradients ,dX,,dγ,.dβ.
,μ=X⊤, the single row as a column, and .σ2=0.Each column has one entry, which is its own mean and has zero deviation from it.
,X^=(X−1μ⊤)D=0⋅D=0, and D=ϵ−1/2I is finite.;X−1μ⊤=X−X=0; with ϵ=0 every entry would be .0/0.
.Y=0⋅diag(γ)+1β⊤=β⊤.1 has one entry.
Y does not depend on X or γ at all, so dX=0 and ,dγ=(dY⊙0)⊤1=0, while .dβ=dY⊤1=dY⊤.Problem 2. Problem 5 agrees: with N=1 and x^=0 it gives .ϵ1(dx^−dx^−0)=0.
With :N=1:,μ=X⊤,,σ2=0,,X^=0,;Y=β⊤;,dX=0,,dγ=0,dβ=dY⊤The layer outputs β whatever the input, so nothing below it receives a gradient and γ 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≥2, one feature takes the same value on every example in the batch. Compute x^ and dx for that feature. What does ϵ do here, and how large is the effect?
μ equals that common value, so x−μ1=0 and .σ2=0.The mean of equal numbers is that number, and every deviation is zero.
.x^=0/0+ϵ=0.ϵ>0 makes the denominator ;ϵ; with ϵ=0 it would be .0/0.
Problem 5 holds at :σ2=0:.dx=Nϵ1(Ndx^−(∑ndx^n)1−0).Its derivation needed only ,σ2+ϵ>0, so the map is smooth here; the last term carries the factor .x^=0.
On a constant feature x^=0 and ;dx=ϵ1(dx^−N1(∑ndx^n)1); without ϵ the forward pass would divide 0 by 0Distribute the N1 over the bracket. The gain is :1/ϵ: about 316 for ,ϵ=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≫ϵ,ϵ changes x^ only by a relative ,ϵ/(2σ2), and the x^ component of dx 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.
X∈RN×d with rows as examplesRight so far: the layout of Problem 1.
“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.
,μ=d1X1∈RN, one mean per exampleThat is layer norm (Problem 7). Batch norm needs ,μ=N1X⊤1∈Rd, one per feature. When N=d the wrong μ has the right length, broadcasts against the rows without error, and subtracts example j's mean from feature .j.
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.
dX^=dYdiag(γ)Right so far: Problem 2.
“Training differs from inference only in which μ and σ2 are used, so the backward pass is Problem 8 with D in place of .Dˉ.”The analogy that causes the mistake: running statistics are constants, but the batch statistics are computed from the current ,X, so they are functions of it.
dX=dYdiag(γ)DIt keeps only the first term of Problem 5 and drops the paths through μ and .σ2. Its column sums are ,(1⊤dX^)D, not 0 (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−1, and a hand-written backward pass often re-derives ∂σ2/∂x from that formula.
dxn=σ2+ϵdx^n+dσ2∂xn∂σ2+NdμRight so far: Problem 5, step 2, with dσ2 and dμ from Problem 4.
“The sample variance is ,N−11∑n(xn−μ)2, so .∂σ2/∂xn=N−12(xn−μ).”The shortcut that causes the mistake: the derivative of the unbiased estimator, while the forward pass divided by N (Problem 3).
dx=Nσ2+ϵ1(Ndx^−(∑ndx^n)1−N−1Nx^∑ndx^nx^n)It is the derivative of a different function from the one the forward pass computed: the x^ term is N/(N−1) times too large. It still sums to zero, so a sum-to-zero test passes, and at batch size 256 the error is under ,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 γ's gradient, and the array of those terms already has a familiar shape.
∂L/∂γj=∑ndYnjX^njRight so far: Problem 2, step 2.
“Each example's contribution is dY⊙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^ is per example, when γ is a parameter shared by all of them.
dγ=dY⊙X^It is ,N×d, not :d: the sum over examples is missing, and the gradient is (dY⊙X^)⊤1 (Problem 2). In array code the update ,γ−ηdY⊙X^, with learning rate ,η, broadcasts silently and turns γ into one scale per example.
5. Scaling the training gradient by the running variance
A batch-norm layer keeps μˉ and σˉ2 as stored attributes, while the batch statistics of the forward pass are temporaries that hand-written code has to remember to save.
In training, X^=(X−1μ⊤)D with the batch statistics μ and σ2Right so far: the training forward pass as defined before Problem 1.
“The layer's variance is ,σˉ2, so use it for the 1/σ2+ϵ in the backward pass.”The shortcut that causes the mistake: reading the stored statistic instead of saving the one the forward pass divided by.
dx=Nσˉ2+ϵ1(Ndx^−(∑ndx^n)1−x^∑ndx^nx^n)The factor comes from differentiating the training forward pass, which divided by σ2+ϵ with the batch variance (Problem 5); σˉ2 is not on the path from X to L in training at all. The answer is off by the factor ,(σ2+ϵ)/(σˉ2+ϵ), close to 1 once the running average has settled, so it hides, and far from 1 early in training, when σˉ2 still holds its initial value.
Print this set: batch-norm-backward.pdf (problems, answers, and worked solutions on separate pages).