Ten problems on the backward pass of layer normalisation: the gradients of the mean and variance, the centering projection, the Jacobian of x̂, the input gradient and why it is orthogonal to 1 and, up to ε, to x̂, the batched version and RMSNorm, with worked solutions and the mistakes that treat μ and σ as constants.
Before you start
Layer normalisation sits in every transformer block: it subtracts each token's mean, divides by its standard deviation, then rescales and shifts with learned γ and .β. The backward pass through γ and β is one line each; the backward pass to the input is where the work is, because the mean and the standard deviation are themselves functions of every input feature. These ten problems build that gradient from its pieces, show why it can only push the input in certain directions, and then batch it and strip it down to RMSNorm. The four mistakes are the ones that give a plausible-looking gradient: μ and σ held constant, γ's gradient read off the output, a 1/σ applied to one term, and a shared parameter's gradient left unsummed.
The conventions are those of the previous pages: vectors are columns, Jacobians are in numerator layout, a gradient has the shape of the variable it is taken with respect to, and for a scalar L of u with u a function of ,x,.∇xL=(∂u/∂x)⊤∇uL.
One token's features are .x∈Rn. Its mean is ,μ=n1∑ixi, its variance ,v=n1∑i(xi−μ)2, and ;σ=v+ϵ; on this page σ is this standard deviation, not the sigmoid.
ϵ is a small constant inside the square root; keep it, it is what stops the Jacobian blowing up on a constant row.
The normalised input is x^=(x−μ1)/σ and the output is ,y=γ⊙x^+β, with learned .γ,β∈Rn.1 is the all-ones vector and ⊙ the elementwise product.
L is a scalar loss that depends on ,x,γ and β only through .y. The upstream gradient is ,g=∇yL, and ;g′=γ⊙g; Problem 5 shows .g′=∇x^L.
,P=I−n111⊤, and mean(u)=n11⊤u for any ,u∈Rn, so .mean(g′⊙x^)=n1x^⊤g′.
Batched: rows are tokens. X∈RN×n and the normalisation is per row: each row gets its own ,μ,v and ,σ, while γ and β are shared by all rows.
Compute .∇xv. Why does the dependence of μ on x not add a term?
··
Let .c=x−μ1. Compute ,∂c/∂x, and show it is symmetric, idempotent and sends 1 to .0.
··
Compute ∇γL and ∇βL in terms of g and .x^.
··
Compute .∇x^L.
···
Compute ∂x^/∂x as a matrix.
···
Compute ∇xL in terms of ,g′=γ⊙g,x^ and ,σ, using only vector operations (no n×n matrix).
··
Show that ,1⊤∇xL=0, and compute ;x^⊤∇xL; when is it ?0? What does that say about which directions of x the loss can push on?
···
Batched: X∈RN×n with rows normalised independently, .G=∇YL. Write ,∇XL,∇γL and .∇βL.
···
RMSNorm drops the mean: y=γ⊙x/r with .r=n1∑ixi2+ϵ. Compute .∇xL.
Answers
∇xμ=n11
∇xv=n2(x−μ1)
;∂c/∂x=I−n111⊤=P;,P⊤=P,,P2=P,P1=0
;∇γL=g⊙x^;∇βL=g
∇x^L=γ⊙g=g′
∂x^/∂x=σ1(I−n111⊤−n1x^x^⊤)
∇xL=σ1(g′−mean(g′)1−x^mean(g′⊙x^))
1⊤∇xL=0 exactly, and ,x^⊤∇xL=σ1v+ϵϵx^⊤g′, which is 0 when ϵ=0 or when ,x^⊤g′=0, and of order ϵ otherwise: the loss cannot move x along 1 (a shift), and can move it along x^ (a rescale) only through the ϵ term
∇XL is Problem 7 applied to each row; ;∇γL=(G⊙X^)⊤1;∇βL=G⊤1
∇xL=r1(g′−x^mean(g′⊙x^)) with x^=x/r
Worked solutions
Problem 1
Compute .∇xμ.
∂μ/∂xj=n1 for every .j.μ=n1∑ixi contains xj once, with coefficient .n1.
∇xμ=n11,n×1, the shape of .x. Every feature moves the mean equally, so in Problem 3 every entry of ∂(μ1)/∂x is .n1.
Problem 2
Compute .∇xv. Why does the dependence of μ on x not add a term?
.∂xj∂v=n1∑i2(xi−μ)(∂xj∂xi−∂xj∂μ).Chain rule on each square, keeping μ as the function of x that it is.
.=n2(xj−μ)−n2∑i(xi−μ)∂xj∂μ.∂xi/∂xj is 1 for i=j and 0 otherwise, which picks out the direct term; ∂μ/∂xj does not depend on ,i, so it can stay inside or come out of the sum.
,∑i(xi−μ)=∑ixi−nμ=0, so the second term vanishes.nμ=∑ixi by the definition of :μ: deviations from the mean sum to zero.
∇xv=n2(x−μ1).n×1. The μ term is not missing by accident: μ is the value of m that minimises ,n1∑i(xi−m)2, so the derivative of that sum with respect to m is zero at ,m=μ, and a small change in μ changes v only to second order.
Problem 3
Let .c=x−μ1. Compute ,∂c/∂x, and show it is symmetric, idempotent and sends 1 to .0.
.∂c/∂x=I−1(∇xμ)⊤.,∂x/∂x=I, and μ1 has entry i equal to ,μ, so row i of its Jacobian is (∇xμ)⊤ for every .i.
,1(∇xμ)⊤=n111⊤, an n×n matrix with every entry .n1.Problem 1.
.P⊤=I⊤−n1(11⊤)⊤=I−n111⊤=P.,(ab⊤)⊤=ba⊤, and here .a=b=1.
.P2=I−n211⊤+n211(1⊤1)1⊤=I−n211⊤+n111⊤=P.Expand the product; ,1⊤1=n, so the last term is .n2n11⊤.
.P1=1−n11(1⊤1)=1−1=0.Again .1⊤1=n.
;∂c/∂x=I−n111⊤=P;,P⊤=P,,P2=P,P1=0.n×n.c is linear in ,x, so :c=Px:P is the orthogonal projection onto the vectors whose entries sum to zero. Centering twice changes nothing (),P2=P), and a constant vector centres to zero ().P1=0).
Problem 4
Compute ∇γL and ∇βL in terms of g and .x^.
.yi=γix^i+βi.⊙ is elementwise, and x^ does not depend on γ or .β.
γi and βi appear only in ,yi, with coefficients x^i and .1.Entry i of γ and of β scales and shifts only feature .i.
∂L/∂γi=gix^i and .∂L/∂βi=gi.L depends on γ and β only through ;y; by step 2 only the yi term of the chain rule is nonzero, and .∂L/∂yi=gi.
;∇γL=g⊙x^;∇βL=gBoth ,n×1, the shapes of γ and .β.
Problem 5
Compute .∇x^L.
,∂y/∂x^=diag(γ),.n×n.yi=γix^i+βi contains only ,x^i, with coefficient .γi.
.∇x^L=diag(γ)⊤g=diag(γ)g.L depends on x^ only through ,y, and a diagonal matrix is its own transpose.
∇x^L=γ⊙g=g′A diagonal matrix times a vector scales entry i by ,γi, so the n×n matrix is never built. From here on the backward pass only needs .g′.
Problem 6
Compute ∂x^/∂x as a matrix.
,x^i=ci/σ, so .∂xj∂x^i=σ1∂xj∂ci−σ2ci∂xj∂σ.Quotient rule, with σ a function of x like .c.
.∇xσ=2σ1∇xv=2σ1⋅n2c=σ1⋅n1(x−μ1)=n1x^.σ=v+ϵ with ϵ constant, so ;dσ/dv=1/(2v+ϵ)=1/(2σ);∇xv is Problem 2. This is where ϵ enters: through σ only.
.∂x^/∂x=σ1P−σ21c(∇xσ)⊤=σ1P−σ1⋅n1x^x^⊤.Step 1 as a matrix: ∂c/∂x=P (Problem 3), and entry (i,j) of c(∇xσ)⊤ is .ci∂σ/∂xj. Then c/σ=x^ and step 2.
∂x^/∂x=σ1(I−n111⊤−n1x^x^⊤)n×n and symmetric. It has the shape of the Jacobians page's Problem 9, the Jacobian of ,x/∥x∥, with a centering added and σ in place of the norm. On a constant row v=0 and ,x^=0, so the Jacobian is :P/ϵ: large but finite, where with ϵ=0 it would divide by zero.
Problem 7
Compute ∇xL in terms of ,g′=γ⊙g,x^ and ,σ, using only vector operations (no n×n matrix).
.∇xL=(∂x^/∂x)⊤g′=(∂x^/∂x)g′.L depends on x only through x^ (through ),y),∇x^L=g′ is Problem 5, and the Jacobian of Problem 6 is symmetric.
.=σ1(g′−n11(1⊤g′)−n1x^(x^⊤g′)).Multiply each term of Problem 6 by g′ and regroup: (11⊤)g′=1(1⊤g′) and ,(x^x^⊤)g′=x^(x^⊤g′), each a vector times a scalar.
n11⊤g′=mean(g′) and .n1x^⊤g′=mean(g′⊙x^).Both are the definition of ;mean; the second uses .x^⊤g′=∑ix^igi′=1⊤(g′⊙x^).
∇xL=σ1(g′−mean(g′)1−x^mean(g′⊙x^)),n×1, the shape of .x. It needs two means and a few elementwise operations, O(n) work, against O(n2) to build and apply the Jacobian. σ and x^ are saved from the forward pass.
Problem 8
Show that ,1⊤∇xL=0, and compute ;x^⊤∇xL; when is it ?0? What does that say about which directions of x the loss can push on?
,1⊤x^=σ11⊤c=σ1∑i(xi−μ)=0, exactly.Deviations from the mean sum to zero (Problem 2, step 3), whatever ϵ is.
.x^⊤x^=σ2c⊤c=v+ϵnv.c⊤c=∑i(xi−μ)2=nv and .σ2=v+ϵ. This equals n only when .ϵ=0.
.1⊤∇xL=σ1(1⊤g′−nmean(g′)−(1⊤x^)mean(g′⊙x^))=σ1(1⊤g′−1⊤g′−0)=0.Problem 7, with ,1⊤1=n,,nmean(g′)=1⊤g′, and step 1.
.x^⊤∇xL=σ1(x^⊤g′−(x^⊤1)mean(g′)−(x^⊤x^)n1x^⊤g′).Problem 7 again, with .mean(g′⊙x^)=n1x^⊤g′.
.=σ1x^⊤g′(1−v+ϵv)=σ1v+ϵϵx^⊤g′.The middle term is 0 by step 1; by step 2 the last is ,v+ϵvx^⊤g′, and .1−v+ϵv=v+ϵϵ. The ϵ survives because ∥x^∥2 falls just short of .n.
1⊤∇xL=0 exactly, and ,x^⊤∇xL=σ1v+ϵϵx^⊤g′, which is 0 when ϵ=0 or when ,x^⊤g′=0, and of order ϵ otherwise: the loss cannot move x along 1 (a shift), and can move it along x^ (a rescale) only through the ϵ termAdding a1 to x adds a to μ and leaves ,c,v and x^ unchanged, so L cannot change. Scaling c by 1+a scales v by ,(1+a)2, and x^=c/v+ϵ would be unchanged if ϵ were .0. With ϵ=10−5 and v near 1 the factor ϵ/(v+ϵ) is about ;10−5; it matters only on rows whose variance is near .ϵ. So a gradient step on x never changes its mean and barely changes its spread: layer norm's output does not see the mean at all, and sees the spread only through .ϵ.
Problem 9
Batched: X∈RN×n with rows normalised independently, .G=∇YL. Write ,∇XL,∇γL and .∇βL.
Write x(k) for token k as a column, so row k of X is ,x(k)⊤, with its own ,μk,,vk,σk and normalised .x^(k).X^ has rows ,x^(k)⊤,Y has rows ,(γ⊙x^(k)+β)⊤,G has rows ,g(k)⊤, and .g′(k)=γ⊙g(k).
Row k of Y depends on X only through row k of .X.Each row's μk and σk are computed from that row alone, and γ and β are not functions of .X.
Row k of ∇XL is .σk1(g′(k)−mean(g′(k))1−x^(k)mean(g′(k)⊙x^(k)))⊤.By step 1 the only path from row k of X to L is through row k of ,Y, whose upstream gradient is ;g(k); the rest is Problem 7 for token .k.
∂L/∂γj=∑kGkjX^kj and .∂L/∂βj=∑kGkj.γj and βj appear in Ykj for every row ,k, with coefficients X^kj and ;1; the chain rule sums over every entry of Y that contains them.
((G⊙X^)⊤1)j=∑k(G⊙X^)kj and ,(G⊤1)j=∑kGkj, with .1∈RN.For a matrix M with N rows, M⊤1 adds up the N rows of ;M; shapes .(n×N)(N×1)=n×1.
∇XL is Problem 7 applied to each row; ;∇γL=(G⊙X^)⊤1;∇βL=G⊤1,N×n,n and ,n, the shapes of ,X,γ and .β. The sums over the N rows appear because γ and β are shared by every token; ∇XL has no sum, because each token has its own row of .X. In array code, the two means in step 2 are taken along the feature axis, one per row.
Problem 10
RMSNorm drops the mean: y=γ⊙x/r with .r=n1∑ixi2+ϵ. Compute .∇xL.
Here write x^=x/r for the RMS-normalised input, and keep .g′=γ⊙g.
.∇x^L=g′.,y=γ⊙x^, so Problem 5 applies unchanged; there is no ,β, and it would not matter if there were.
.∇xr=2r1⋅n2x=n1x^.,r=n1∑ixi2+ϵ, the derivative of n1∑ixi2 with respect to xj is ,n2xj, and .x/r=x^. Problem 6, step 2, with x in place of .c.
.∂x^/∂x=r1I−r21x(∇xr)⊤=r1(I−n1x^x^⊤).Problem 6, step 3, with ∂x/∂x=I in place of :∂c/∂x=P: nothing is subtracted, so there is no centering term.
.∇xL=r1(g′−n1x^(x^⊤g′)).The Jacobian is symmetric, so apply it to g′ as in Problem 7.
∇xL=r1(g′−x^mean(g′⊙x^)) with x^=x/r.n×1. It is Problem 7 without the mean(g′)1 term. Because RMSNorm does not subtract the mean, a shift of x does change its output, and 1⊤∇xL is not zero in general; only the rescale direction is (nearly) invisible, as in Problem 8.
Where this goes wrong
1. Treating μ and σ as constants
Per feature, layer norm looks like an affine map, ,xi↦γi(xi−μ)/σ+βi, and an affine map's input gradient is just its slope times the upstream gradient.
∇x^L=g′Right so far: Problem 5.
“x^=(x−μ1)/σ subtracts a number and divides by a number, so .∂x^/∂x=I/σ.”The shortcut that causes the mistake: μ and σ are treated as numbers fixed in the forward pass, when both are functions of every .xi.
∇xL=g′/σThat is the Jacobian's first term only (Problem 6); the other two come from ∇xμ and ,∇xσ, and the centering term is what makes 1⊤∇xL=0 (Problem 8). This answer sums to ,1⊤g′/σ, so it claims that a uniform shift of x changes the loss, which it cannot.
2. Gradient of γ from the output instead of the normalised input
In a linear layer a weight's gradient is the upstream gradient times the layer's input, and it is easy to reach for the nearest saved vector instead.
yi=γix^i+βi and ∂L/∂yi=giRight so far: the forward pass and the upstream gradient.
“The gradient of a scale is the upstream gradient times what it scales, and what comes out is .y.”The shortcut that causes the mistake: using the saved output of the layer in place of the input that γ multiplies.
∇γL=g⊙yy already contains :γ:.g⊙y=γ⊙g⊙x^+g⊙β. The coefficient of γi in yi is ,x^i, so the answer is g⊙x^ (Problem 4). The two agree at the usual initialisation, γ=1 and ,β=0, so the bug passes a gradient check at step 0.
3. Dividing only the first term by σ
Written out term by term, Problem 7's leading σ1 is easy to attach to the first term only.
∂x^/∂x=σ1(I−n111⊤−n1x^x^⊤)Right so far: Problem 6.
“Apply it to :g′:,g′/σ, then subtract the mean of g′ and the x^ term.”The slip that causes the mistake: scaling only the identity term, as if σ1 belonged to .I.
∇xL=g′/σ−mean(g′)1−x^mean(g′⊙x^)The σ1 multiplies the whole Jacobian, so every term. This answer sums to ,(σ1−1)1⊤g′, not ,0, and is right only when ,σ=1, so it looks right on rows already close to unit variance.
4. Forgetting to sum γ's gradient over the batch
For one token ∇γL=g⊙x^ (Problem 4); in a batch g and x^ become matrices.
∂L/∂γj collects a term GkjX^kj from every row kRight so far: Problem 9, step 3.
“Replace the vectors by their batched matrices.”The shortcut that causes the mistake: ∇XL is the per-token formula row by row, and the same substitution is applied to a parameter.
∇γL=G⊙X^γ is shared by every row, so its gradient is the sum of the per-row gradients, (G⊙X^)⊤1 (Problem 9): length ,n, not .N×n. In array code the update ,γ−ηG⊙X^, with learning rate ,η, broadcasts silently and turns γ into one scale per token.
Print this set: layer-norm-backward.pdf (problems, answers, and worked solutions on separate pages).