Practice / Transformer pieces

Layer norm backward

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 γ\gamma and β\beta. The backward pass through γ\gamma and β\beta 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: μ\mu and σ\sigma held constant, γ\gamma's gradient read off the output, a 1/σ1/\sigma 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 LL of uu with uu a function of xx, ∇xL=(∂u/∂x)⊤∇uL\nabla_x L = (\partial u/\partial x)^\top \nabla_u L.
  • One token's features are x∈Rnx \in \mathbb{R}^n. Its mean is μ=1n∑ixi\mu = \tfrac1n\sum_i x_i, its variance v=1n∑i(xi−μ)2v = \tfrac1n\sum_i (x_i-\mu)^2, and σ=v+ϵ\sigma = \sqrt{v + \epsilon}; on this page σ\sigma is this standard deviation, not the sigmoid.
  • ϵ\epsilon 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)/σ\hat x = (x - \mu\mathbf{1})/\sigma and the output is y=γ⊙x^+βy = \gamma\odot\hat x + \beta, with learned γ,β∈Rn\gamma, \beta \in \mathbb{R}^n. 1\mathbf{1} is the all-ones vector and ⊙\odot the elementwise product.
  • LL is a scalar loss that depends on xx, γ\gamma and β\beta only through yy. The upstream gradient is g=∇yLg = \nabla_y L, and g′=γ⊙gg' = \gamma\odot g; Problem 5 shows g′=∇x^Lg' = \nabla_{\hat x} L.
  • P=I−1n11⊤P = I - \tfrac1n\mathbf{1}\mathbf{1}^\top, and mean⁡(u)=1n1⊤u\operatorname{mean}(u) = \tfrac1n\mathbf{1}^\top u for any u∈Rnu \in \mathbb{R}^n, so mean⁡(g′⊙x^)=1nx^⊤g′\operatorname{mean}(g'\odot\hat x) = \tfrac1n\hat x^\top g'.
  • Batched: rows are tokens. X∈RN×nX \in \mathbb{R}^{N\times n} and the normalisation is per row: each row gets its own μ\mu, vv and σ\sigma, while γ\gamma and β\beta are shared by all rows.

Builds on: Matrix calculus conventions, Jacobians and the chain rule

Problems

  1. ·

    Compute ∇xμ\nabla_x \mu.

  2. ··

    Compute ∇xv\nabla_x v. Why does the dependence of μ\mu on xx not add a term?

  3. ··

    Let c=x−μ1c = x - \mu\mathbf{1}. Compute ∂c/∂x\partial c/\partial x, and show it is symmetric, idempotent and sends 1\mathbf{1} to 00.

  4. ··

    Compute ∇γL\nabla_\gamma L and ∇βL\nabla_\beta L in terms of gg and x^\hat x.

  5. ··

    Compute ∇x^L\nabla_{\hat x} L.

  6. ···

    Compute ∂x^/∂x\partial\hat x/\partial x as a matrix.

  7. ···

    Compute ∇xL\nabla_x L in terms of g′=γ⊙gg' = \gamma\odot g, x^\hat x and σ\sigma, using only vector operations (no n×nn\times n matrix).

  8. ··

    Show that 1⊤∇xL=0\mathbf{1}^\top\nabla_x L = 0, and compute x^⊤∇xL\hat x^\top\nabla_x L; when is it 00? What does that say about which directions of xx the loss can push on?

  9. ···

    Batched: X∈RN×nX \in \mathbb{R}^{N\times n} with rows normalised independently, G=∇YLG = \nabla_Y L. Write ∇XL\nabla_X L, ∇γL\nabla_\gamma L and ∇βL\nabla_\beta L.

  10. ···

    RMSNorm drops the mean: y=γ⊙x/ry = \gamma\odot x/r with r=1n∑ixi2+ϵr = \sqrt{\tfrac1n\sum_i x_i^2 + \epsilon}. Compute ∇xL\nabla_x L.

Worked solutions

Problem 1

Compute ∇xμ\nabla_x \mu.

  1. ∂μ/∂xj=1n\partial\mu/\partial x_j = \tfrac1n for every jj.μ=1n∑ixi\mu = \tfrac1n\sum_i x_i contains xjx_j once, with coefficient 1n\tfrac1n.
  2. ∇xμ=1n1\nabla_x\mu = \tfrac1n\mathbf{1}n×1n \times 1, the shape of xx. Every feature moves the mean equally, so in Problem 3 every entry of ∂(μ1)/∂x\partial(\mu\mathbf{1})/\partial x is 1n\tfrac1n.

Problem 2

Compute ∇xv\nabla_x v. Why does the dependence of μ\mu on xx not add a term?

  1. ∂v∂xj=1n∑i2(xi−μ)(∂xi∂xj−∂μ∂xj)\dfrac{\partial v}{\partial x_j} = \dfrac1n\sum_i 2(x_i-\mu)\Big(\dfrac{\partial x_i}{\partial x_j} - \dfrac{\partial\mu}{\partial x_j}\Big).Chain rule on each square, keeping μ\mu as the function of xx that it is.
  2. =2n(xj−μ)−2n∑i(xi−μ) ∂μ∂xj= \tfrac2n(x_j-\mu) - \tfrac2n\sum_i (x_i-\mu)\,\dfrac{\partial\mu}{\partial x_j}.∂xi/∂xj\partial x_i/\partial x_j is 11 for i=ji = j and 00 otherwise, which picks out the direct term; ∂μ/∂xj\partial\mu/\partial x_j does not depend on ii, so it can stay inside or come out of the sum.
  3. ∑i(xi−μ)=∑ixi−nμ=0\sum_i (x_i-\mu) = \sum_i x_i - n\mu = 0, so the second term vanishes.nμ=∑ixin\mu = \sum_i x_i by the definition of μ\mu: deviations from the mean sum to zero.
  4. ∇xv=2n(x−μ1)\nabla_x v = \tfrac2n(x - \mu\mathbf{1})n×1n \times 1. The μ\mu term is not missing by accident: μ\mu is the value of mm that minimises 1n∑i(xi−m)2\tfrac1n\sum_i (x_i - m)^2, so the derivative of that sum with respect to mm is zero at m=μm = \mu, and a small change in μ\mu changes vv only to second order.

Problem 3

Let c=x−μ1c = x - \mu\mathbf{1}. Compute ∂c/∂x\partial c/\partial x, and show it is symmetric, idempotent and sends 1\mathbf{1} to 00.

  1. ∂c/∂x=I−1 (∇xμ)⊤\partial c/\partial x = I - \mathbf{1}\,(\nabla_x\mu)^\top.∂x/∂x=I\partial x/\partial x = I, and μ1\mu\mathbf{1} has entry ii equal to μ\mu, so row ii of its Jacobian is (∇xμ)⊤(\nabla_x\mu)^\top for every ii.
  2. 1 (∇xμ)⊤=1n11⊤\mathbf{1}\,(\nabla_x\mu)^\top = \tfrac1n\mathbf{1}\mathbf{1}^\top, an n×nn \times n matrix with every entry 1n\tfrac1n.Problem 1.
  3. P⊤=I⊤−1n(11⊤)⊤=I−1n11⊤=PP^\top = I^\top - \tfrac1n(\mathbf{1}\mathbf{1}^\top)^\top = I - \tfrac1n\mathbf{1}\mathbf{1}^\top = P.(ab⊤)⊤=ba⊤(ab^\top)^\top = ba^\top, and here a=b=1a = b = \mathbf{1}.
  4. P2=I−2n11⊤+1n21(1⊤1)1⊤=I−2n11⊤+1n11⊤=PP^2 = I - \tfrac2n\mathbf{1}\mathbf{1}^\top + \tfrac1{n^2}\mathbf{1}(\mathbf{1}^\top\mathbf{1})\mathbf{1}^\top = I - \tfrac2n\mathbf{1}\mathbf{1}^\top + \tfrac1n\mathbf{1}\mathbf{1}^\top = P.Expand the product; 1⊤1=n\mathbf{1}^\top\mathbf{1} = n, so the last term is nn211⊤\tfrac{n}{n^2}\mathbf{1}\mathbf{1}^\top.
  5. P1=1−1n1(1⊤1)=1−1=0P\mathbf{1} = \mathbf{1} - \tfrac1n\mathbf{1}(\mathbf{1}^\top\mathbf{1}) = \mathbf{1} - \mathbf{1} = 0.Again 1⊤1=n\mathbf{1}^\top\mathbf{1} = n.
  6. ∂c/∂x=I−1n11⊤=P\partial c/\partial x = I - \tfrac1n\mathbf{1}\mathbf{1}^\top = P; P⊤=PP^\top = P, P2=PP^2 = P, P1=0P\mathbf{1} = 0n×nn \times n. cc is linear in xx, so c=Pxc = Px: PP is the orthogonal projection onto the vectors whose entries sum to zero. Centering twice changes nothing (P2=PP^2 = P), and a constant vector centres to zero (P1=0P\mathbf{1} = 0).

Problem 4

Compute ∇γL\nabla_\gamma L and ∇βL\nabla_\beta L in terms of gg and x^\hat x.

  1. yi=γix^i+βiy_i = \gamma_i\hat x_i + \beta_i.⊙\odot is elementwise, and x^\hat x does not depend on γ\gamma or β\beta.
  2. γi\gamma_i and βi\beta_i appear only in yiy_i, with coefficients x^i\hat x_i and 11.Entry ii of γ\gamma and of β\beta scales and shifts only feature ii.
  3. ∂L/∂γi=gix^i\partial L/\partial\gamma_i = g_i\hat x_i and ∂L/∂βi=gi\partial L/\partial\beta_i = g_i.LL depends on γ\gamma and β\beta only through yy; by step 2 only the yiy_i term of the chain rule is nonzero, and ∂L/∂yi=gi\partial L/\partial y_i = g_i.
  4. ∇γL=g⊙x^\nabla_\gamma L = g\odot\hat x; ∇βL=g\nabla_\beta L = gBoth n×1n \times 1, the shapes of γ\gamma and β\beta.

Problem 5

Compute ∇x^L\nabla_{\hat x} L.

  1. ∂y/∂x^=diag⁡(γ)\partial y/\partial\hat x = \operatorname{diag}(\gamma), n×nn \times n.yi=γix^i+βiy_i = \gamma_i\hat x_i + \beta_i contains only x^i\hat x_i, with coefficient γi\gamma_i.
  2. ∇x^L=diag⁡(γ)⊤g=diag⁡(γ) g\nabla_{\hat x} L = \operatorname{diag}(\gamma)^\top g = \operatorname{diag}(\gamma)\,g.LL depends on x^\hat x only through yy, and a diagonal matrix is its own transpose.
  3. ∇x^L=γ⊙g=g′\nabla_{\hat x} L = \gamma\odot g = g'A diagonal matrix times a vector scales entry ii by γi\gamma_i, so the n×nn \times n matrix is never built. From here on the backward pass only needs g′g'.

Problem 6

Compute ∂x^/∂x\partial\hat x/\partial x as a matrix.

  1. x^i=ci/σ\hat x_i = c_i/\sigma, so ∂x^i∂xj=1σ∂ci∂xj−ciσ2∂σ∂xj\dfrac{\partial\hat x_i}{\partial x_j} = \dfrac1\sigma\dfrac{\partial c_i}{\partial x_j} - \dfrac{c_i}{\sigma^2}\dfrac{\partial\sigma}{\partial x_j}.Quotient rule, with σ\sigma a function of xx like cc.
  2. ∇xσ=12σ∇xv=12σ⋅2n c=1σ⋅1n(x−μ1)=1nx^\nabla_x\sigma = \dfrac{1}{2\sigma}\nabla_x v = \dfrac1{2\sigma}\cdot\dfrac2n\,c = \dfrac1\sigma\cdot\dfrac1n(x-\mu\mathbf{1}) = \dfrac1n\hat x.σ=v+ϵ\sigma = \sqrt{v+\epsilon} with ϵ\epsilon constant, so dσ/dv=1/(2v+ϵ)=1/(2σ)d\sigma/dv = 1/(2\sqrt{v+\epsilon}) = 1/(2\sigma); ∇xv\nabla_x v is Problem 2. This is where ϵ\epsilon enters: through σ\sigma only.
  3. ∂x^/∂x=1σP−1σ2 c (∇xσ)⊤=1σP−1σ⋅1n x^x^⊤\partial\hat x/\partial x = \dfrac1\sigma P - \dfrac1{\sigma^2}\,c\,(\nabla_x\sigma)^\top = \dfrac1\sigma P - \dfrac1\sigma\cdot\dfrac1n\,\hat x\hat x^\top.Step 1 as a matrix: ∂c/∂x=P\partial c/\partial x = P (Problem 3), and entry (i,j)(i,j) of c (∇xσ)⊤c\,(\nabla_x\sigma)^\top is ci ∂σ/∂xjc_i\,\partial\sigma/\partial x_j. Then c/σ=x^c/\sigma = \hat x and step 2.
  4. ∂x^/∂x=1σ(I−1n11⊤−1nx^x^⊤)\partial\hat x/\partial x = \tfrac1\sigma\big(I - \tfrac1n\mathbf{1}\mathbf{1}^\top - \tfrac1n\hat x\hat x^\top\big)n×nn \times n and symmetric. It has the shape of the Jacobians page's Problem 9, the Jacobian of x/∥x∥x/\|x\|, with a centering added and σ\sigma in place of the norm. On a constant row v=0v = 0 and x^=0\hat x = 0, so the Jacobian is P/ϵP/\sqrt\epsilon: large but finite, where with ϵ=0\epsilon = 0 it would divide by zero.

Problem 7

Compute ∇xL\nabla_x L in terms of g′=γ⊙gg' = \gamma\odot g, x^\hat x and σ\sigma, using only vector operations (no n×nn\times n matrix).

  1. ∇xL=(∂x^/∂x)⊤g′=(∂x^/∂x) g′\nabla_x L = (\partial\hat x/\partial x)^\top g' = (\partial\hat x/\partial x)\,g'.LL depends on xx only through x^\hat x (through yy), ∇x^L=g′\nabla_{\hat x} L = g' is Problem 5, and the Jacobian of Problem 6 is symmetric.
  2. =1σ(g′−1n1(1⊤g′)−1nx^(x^⊤g′))= \tfrac1\sigma\big(g' - \tfrac1n\mathbf{1}(\mathbf{1}^\top g') - \tfrac1n\hat x(\hat x^\top g')\big).Multiply each term of Problem 6 by g′g' and regroup: (11⊤)g′=1(1⊤g′)(\mathbf{1}\mathbf{1}^\top)g' = \mathbf{1}(\mathbf{1}^\top g') and (x^x^⊤)g′=x^(x^⊤g′)(\hat x\hat x^\top)g' = \hat x(\hat x^\top g'), each a vector times a scalar.
  3. 1n1⊤g′=mean⁡(g′)\tfrac1n\mathbf{1}^\top g' = \operatorname{mean}(g') and 1nx^⊤g′=mean⁡(g′⊙x^)\tfrac1n\hat x^\top g' = \operatorname{mean}(g'\odot\hat x).Both are the definition of mean⁡\operatorname{mean}; the second uses x^⊤g′=∑ix^igi′=1⊤(g′⊙x^)\hat x^\top g' = \sum_i \hat x_i g'_i = \mathbf{1}^\top(g'\odot\hat x).
  4. ∇xL=1σ(g′−mean⁡(g′)1−x^mean⁡(g′⊙x^))\nabla_x L = \tfrac1\sigma\big(g' - \operatorname{mean}(g')\mathbf{1} - \hat x\operatorname{mean}(g'\odot\hat x)\big)n×1n \times 1, the shape of xx. It needs two means and a few elementwise operations, O(n)O(n) work, against O(n2)O(n^2) to build and apply the Jacobian. σ\sigma and x^\hat x are saved from the forward pass.

Problem 8

Show that 1⊤∇xL=0\mathbf{1}^\top\nabla_x L = 0, and compute x^⊤∇xL\hat x^\top\nabla_x L; when is it 00? What does that say about which directions of xx the loss can push on?

  1. 1⊤x^=1σ1⊤c=1σ∑i(xi−μ)=0\mathbf{1}^\top\hat x = \tfrac1\sigma\mathbf{1}^\top c = \tfrac1\sigma\sum_i (x_i - \mu) = 0, exactly.Deviations from the mean sum to zero (Problem 2, step 3), whatever ϵ\epsilon is.
  2. x^⊤x^=c⊤cσ2=nvv+ϵ\hat x^\top\hat x = \dfrac{c^\top c}{\sigma^2} = \dfrac{nv}{v+\epsilon}.c⊤c=∑i(xi−μ)2=nvc^\top c = \sum_i (x_i-\mu)^2 = nv and σ2=v+ϵ\sigma^2 = v + \epsilon. This equals nn only when ϵ=0\epsilon = 0.
  3. 1⊤∇xL=1σ(1⊤g′−nmean⁡(g′)−(1⊤x^)mean⁡(g′⊙x^))=1σ(1⊤g′−1⊤g′−0)=0\mathbf{1}^\top\nabla_x L = \tfrac1\sigma\big(\mathbf{1}^\top g' - n\operatorname{mean}(g') - (\mathbf{1}^\top\hat x)\operatorname{mean}(g'\odot\hat x)\big) = \tfrac1\sigma(\mathbf{1}^\top g' - \mathbf{1}^\top g' - 0) = 0.Problem 7, with 1⊤1=n\mathbf{1}^\top\mathbf{1} = n, nmean⁡(g′)=1⊤g′n\operatorname{mean}(g') = \mathbf{1}^\top g', and step 1.
  4. x^⊤∇xL=1σ(x^⊤g′−(x^⊤1)mean⁡(g′)−(x^⊤x^) 1nx^⊤g′)\hat x^\top\nabla_x L = \tfrac1\sigma\big(\hat x^\top g' - (\hat x^\top\mathbf{1})\operatorname{mean}(g') - (\hat x^\top\hat x)\,\tfrac1n\hat x^\top g'\big).Problem 7 again, with mean⁡(g′⊙x^)=1nx^⊤g′\operatorname{mean}(g'\odot\hat x) = \tfrac1n\hat x^\top g'.
  5. =1σ x^⊤g′(1−vv+ϵ)=1σ ϵv+ϵ x^⊤g′= \tfrac1\sigma\,\hat x^\top g'\Big(1 - \dfrac{v}{v+\epsilon}\Big) = \dfrac1\sigma\,\dfrac{\epsilon}{v+\epsilon}\,\hat x^\top g'.The middle term is 00 by step 1; by step 2 the last is vv+ϵx^⊤g′\tfrac{v}{v+\epsilon}\hat x^\top g', and 1−vv+ϵ=ϵv+ϵ1 - \tfrac{v}{v+\epsilon} = \tfrac{\epsilon}{v+\epsilon}. The ϵ\epsilon survives because ∥x^∥2\|\hat x\|^2 falls just short of nn.
  6. 1⊤∇xL=0\mathbf{1}^\top\nabla_x L = 0 exactly, and x^⊤∇xL=1σ ϵv+ϵ x^⊤g′\hat x^\top\nabla_x L = \dfrac{1}{\sigma}\,\dfrac{\epsilon}{v+\epsilon}\,\hat x^\top g', which is 00 when ϵ=0\epsilon = 0 or when x^⊤g′=0\hat x^\top g' = 0, and of order ϵ\epsilon otherwise: the loss cannot move xx along 1\mathbf{1} (a shift), and can move it along x^\hat x (a rescale) only through the ϵ\epsilon termAdding a1a\mathbf{1} to xx adds aa to μ\mu and leaves cc, vv and x^\hat x unchanged, so LL cannot change. Scaling cc by 1+a1 + a scales vv by (1+a)2(1+a)^2, and x^=c/v+ϵ\hat x = c/\sqrt{v+\epsilon} would be unchanged if ϵ\epsilon were 00. With ϵ=10−5\epsilon = 10^{-5} and vv near 11 the factor ϵ/(v+ϵ)\epsilon/(v+\epsilon) is about 10−510^{-5}; it matters only on rows whose variance is near ϵ\epsilon. So a gradient step on xx 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 ϵ\epsilon.

Problem 9

Batched: X∈RN×nX \in \mathbb{R}^{N\times n} with rows normalised independently, G=∇YLG = \nabla_Y L. Write ∇XL\nabla_X L, ∇γL\nabla_\gamma L and ∇βL\nabla_\beta L.

Write x(k)x^{(k)} for token kk as a column, so row kk of XX is x(k)⊤x^{(k)\top}, with its own μk\mu_k, vkv_k, σk\sigma_k and normalised x^(k)\hat x^{(k)}. X^\hat X has rows x^(k)⊤\hat x^{(k)\top}, YY has rows (γ⊙x^(k)+β)⊤(\gamma\odot\hat x^{(k)} + \beta)^\top, GG has rows g(k)⊤g^{(k)\top}, and g′(k)=γ⊙g(k)g'^{(k)} = \gamma\odot g^{(k)}.

  1. Row kk of YY depends on XX only through row kk of XX.Each row's μk\mu_k and σk\sigma_k are computed from that row alone, and γ\gamma and β\beta are not functions of XX.
  2. Row kk of ∇XL\nabla_X L is 1σk(g′(k)−mean⁡(g′(k))1−x^(k)mean⁡(g′(k)⊙x^(k)))⊤\tfrac1{\sigma_k}\big(g'^{(k)} - \operatorname{mean}(g'^{(k)})\mathbf{1} - \hat x^{(k)}\operatorname{mean}(g'^{(k)}\odot\hat x^{(k)})\big)^\top.By step 1 the only path from row kk of XX to LL is through row kk of YY, whose upstream gradient is g(k)g^{(k)}; the rest is Problem 7 for token kk.
  3. ∂L/∂γj=∑kGkjX^kj\partial L/\partial\gamma_j = \sum_k G_{kj}\hat X_{kj} and ∂L/∂βj=∑kGkj\partial L/\partial\beta_j = \sum_k G_{kj}.γj\gamma_j and βj\beta_j appear in YkjY_{kj} for every row kk, with coefficients X^kj\hat X_{kj} and 11; the chain rule sums over every entry of YY that contains them.
  4. ((G⊙X^)⊤1)j=∑k(G⊙X^)kj\big((G\odot\hat X)^\top\mathbf{1}\big)_j = \sum_k (G\odot\hat X)_{kj} and (G⊤1)j=∑kGkj(G^\top\mathbf{1})_j = \sum_k G_{kj}, with 1∈RN\mathbf{1} \in \mathbb{R}^N.For a matrix MM with NN rows, M⊤1M^\top\mathbf{1} adds up the NN rows of MM; shapes (n×N)(N×1)=n×1(n\times N)(N\times 1) = n \times 1.
  5. ∇XL\nabla_X L is Problem 7 applied to each row; ∇γL=(G⊙X^)⊤1\nabla_\gamma L = (G\odot\hat X)^\top\mathbf{1}; ∇βL=G⊤1\nabla_\beta L = G^\top\mathbf{1}N×nN \times n, nn and nn, the shapes of XX, γ\gamma and β\beta. The sums over the NN rows appear because γ\gamma and β\beta are shared by every token; ∇XL\nabla_X L has no sum, because each token has its own row of XX. 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/ry = \gamma\odot x/r with r=1n∑ixi2+ϵr = \sqrt{\tfrac1n\sum_i x_i^2 + \epsilon}. Compute ∇xL\nabla_x L.

Here write x^=x/r\hat x = x/r for the RMS-normalised input, and keep g′=γ⊙gg' = \gamma\odot g.

  1. ∇x^L=g′\nabla_{\hat x} L = g'.y=γ⊙x^y = \gamma\odot\hat x, so Problem 5 applies unchanged; there is no β\beta, and it would not matter if there were.
  2. ∇xr=12r⋅2n x=1nx^\nabla_x r = \dfrac1{2r}\cdot\dfrac2n\,x = \dfrac1n\hat x.r=1n∑ixi2+ϵr = \sqrt{\tfrac1n\sum_i x_i^2 + \epsilon}, the derivative of 1n∑ixi2\tfrac1n\sum_i x_i^2 with respect to xjx_j is 2nxj\tfrac2n x_j, and x/r=x^x/r = \hat x. Problem 6, step 2, with xx in place of cc.
  3. ∂x^/∂x=1rI−1r2 x (∇xr)⊤=1r(I−1nx^x^⊤)\partial\hat x/\partial x = \tfrac1r I - \tfrac1{r^2}\,x\,(\nabla_x r)^\top = \tfrac1r\big(I - \tfrac1n\hat x\hat x^\top\big).Problem 6, step 3, with ∂x/∂x=I\partial x/\partial x = I in place of ∂c/∂x=P\partial c/\partial x = P: nothing is subtracted, so there is no centering term.
  4. ∇xL=1r(g′−1nx^(x^⊤g′))\nabla_x L = \tfrac1r\big(g' - \tfrac1n\hat x(\hat x^\top g')\big).The Jacobian is symmetric, so apply it to g′g' as in Problem 7.
  5. ∇xL=1r(g′−x^mean⁡(g′⊙x^))\nabla_x L = \tfrac1r\big(g' - \hat x\operatorname{mean}(g'\odot\hat x)\big) with x^=x/r\hat x = x/rn×1n \times 1. It is Problem 7 without the mean⁡(g′)1\operatorname{mean}(g')\mathbf{1} term. Because RMSNorm does not subtract the mean, a shift of xx does change its output, and 1⊤∇xL\mathbf{1}^\top\nabla_x L 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−μ)/σ+βix_i \mapsto \gamma_i(x_i - \mu)/\sigma + \beta_i, and an affine map's input gradient is just its slope times the upstream gradient.

  1. ∇x^L=g′\nabla_{\hat x} L = g'Right so far: Problem 5.
  2. “x^=(x−μ1)/σ\hat x = (x - \mu\mathbf{1})/\sigma subtracts a number and divides by a number, so ∂x^/∂x=I/σ\partial\hat x/\partial x = I/\sigma.”The shortcut that causes the mistake: μ\mu and σ\sigma are treated as numbers fixed in the forward pass, when both are functions of every xix_i.
  3. ∇xL=g′/σ\nabla_x L = g'/\sigmaThat is the Jacobian's first term only (Problem 6); the other two come from ∇xμ\nabla_x\mu and ∇xσ\nabla_x\sigma, and the centering term is what makes 1⊤∇xL=0\mathbf{1}^\top\nabla_x L = 0 (Problem 8). This answer sums to 1⊤g′/σ\mathbf{1}^\top g'/\sigma, so it claims that a uniform shift of xx 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.

  1. yi=γix^i+βiy_i = \gamma_i\hat x_i + \beta_i and ∂L/∂yi=gi\partial L/\partial y_i = g_iRight so far: the forward pass and the upstream gradient.
  2. “The gradient of a scale is the upstream gradient times what it scales, and what comes out is yy.”The shortcut that causes the mistake: using the saved output of the layer in place of the input that γ\gamma multiplies.
  3. ∇γL=g⊙y\nabla_\gamma L = g\odot yyy already contains γ\gamma: g⊙y=γ⊙g⊙x^+g⊙βg\odot y = \gamma\odot g\odot\hat x + g\odot\beta. The coefficient of γi\gamma_i in yiy_i is x^i\hat x_i, so the answer is g⊙x^g\odot\hat x (Problem 4). The two agree at the usual initialisation, γ=1\gamma = \mathbf{1} and β=0\beta = 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σ\tfrac1\sigma is easy to attach to the first term only.

  1. ∂x^/∂x=1σ(I−1n11⊤−1nx^x^⊤)\partial\hat x/\partial x = \tfrac1\sigma\big(I - \tfrac1n\mathbf{1}\mathbf{1}^\top - \tfrac1n\hat x\hat x^\top\big)Right so far: Problem 6.
  2. “Apply it to g′g': g′/σg'/\sigma, then subtract the mean of g′g' and the x^\hat x term.”The slip that causes the mistake: scaling only the identity term, as if 1σ\tfrac1\sigma belonged to II.
  3. ∇xL=g′/σ−mean⁡(g′)1−x^mean⁡(g′⊙x^)\nabla_x L = g'/\sigma - \operatorname{mean}(g')\mathbf{1} - \hat x\operatorname{mean}(g'\odot\hat x)The 1σ\tfrac1\sigma multiplies the whole Jacobian, so every term. This answer sums to (1σ−1)1⊤g′(\tfrac1\sigma - 1)\mathbf{1}^\top g', not 00, and is right only when σ=1\sigma = 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^\nabla_\gamma L = g\odot\hat x (Problem 4); in a batch gg and x^\hat x become matrices.

  1. ∂L/∂γj\partial L/\partial\gamma_j collects a term GkjX^kjG_{kj}\hat X_{kj} from every row kkRight so far: Problem 9, step 3.
  2. “Replace the vectors by their batched matrices.”The shortcut that causes the mistake: ∇XL\nabla_X L is the per-token formula row by row, and the same substitution is applied to a parameter.
  3. ∇γL=G⊙X^\nabla_\gamma L = G\odot\hat Xγ\gamma is shared by every row, so its gradient is the sum of the per-row gradients, (G⊙X^)⊤1(G\odot\hat X)^\top\mathbf{1} (Problem 9): length nn, not N×nN\times n. In array code the update γ−η G⊙X^\gamma - \eta\,G\odot\hat X, with learning rate η\eta, broadcasts silently and turns γ\gamma into one scale per token.

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