Skip to content

Matrix differentials, the trace trick, and closed-form VJPs

ModuleM08.3 · build · Python · Pass 2 · 3 to 4 h
You buildpython/tinyllm/autograd/vjp.py: unbroadcast, matmul_vjp, softmax_vjp, log_softmax_vjp, layernorm_vjp, rmsnorm_vjp, cross_entropy_vjp
Contractcourse/contracts/py/tinyllm/autograd/vjp.pyi
Testscourse/tests/M08.3/ (what they check: section 4)
NeedsM09.2 stable softmax · M11.1 cross-entropy (the tests differentiate it) · M04.2 numeric VJP · reading: M03.1 matrices and matmul shapes (or --ref-deps)
Used byL0.1 unbroadcast in your Tensor’s backward · L0.2 op VJPs · L0.3 fused cross-entropy · later: L3.1 backpropagation through time by hand, L7.1 RMSNorm
MilestoneMS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains)
Optional depthParr and Howard, The Matrix Calculus You Need for Deep Learning; Minka, “Old and New Matrix Algebra Useful for Statistics” (the differential method); Petersen and Pedersen, The Matrix Cookbook, sections 2 and 4
  • For a scalar loss, dL=tr⁡(G⊤dY)dL = \operatorname{tr}(G^\top dY) with G=∂L/∂YG = \partial L / \partial Y; write dYdY in terms of dXdX, move dXdX to the right with the cyclic property of the trace, and the matrix in front of it is ∂L/∂X\partial L / \partial X (test_trace_identity).
  • That gives every rule in a few lines: Aˉ=GB⊤\bar A = G B^\top and Bˉ=A⊤G\bar B = A^\top G for a matmul, y⊙(g−⟨g,y⟩)y \odot (g - \langle g, y\rangle) for softmax, (softmax−onehot)/n(\mathrm{softmax} - \mathrm{onehot}) / n for cross-entropy (test_hand_example, test_matmul_vjp_gradcheck).
  • LayerNorm and RMSNorm divide by a statistic of every input, so their VJPs carry correction terms that a “treat the statistic as a constant” derivation drops (test_layernorm_vjp_gradcheck, test_rmsnorm_vjp_gradcheck).
  • A broadcast input receives the sum of its copies’ gradients (test_unbroadcast_shapes), and padded positions receive exactly zero (test_cross_entropy_ignore_index).
Terminal window
ol start M08.3 # stubs vjp.py into your repo, contract alongside
ol tests M08.3 # read the test catalog first: rung R0, you write no tests here
ol check M08.3 # exit code is the verdict
ol check M08.3 --ref-deps # only if your M09.2, M11.1, or M04.2 is not passing yet
ol diff M08.3 # after passing: your code against the reference

M08.2 gave you reverse mode one scalar at a time. Your bigram’s forward pass is a [T,256]×[256,256][T, 256] \times [256, 256] matmul followed by a softmax over 256 entries per row, about 16 million multiply-adds per step; spelled out as Python Value objects that is minutes per step instead of milliseconds. L0.2 builds a tensor op library instead, where each op (matmul, softmax, log-softmax, LayerNorm, RMSNorm, cross-entropy) has one backward function written in numpy. Each needs a closed-form vector-Jacobian product, derived once on paper and proven by gradcheck. This module derives and implements them, so that L0.2 only wires them into the graph, L0.3 fuses softmax with cross-entropy, L3.1 reuses them to backpropagate through time by hand, and L7.1’s RMSNorm has its backward ready.

SymbolMeaningType / shape
LLa scalar lossfloat
X,Y=f(X)X, Y = f(X)an op’s input and outputarrays
G=Yˉ=∂L/∂YG = \bar Y = \partial L / \partial Ythe upstream gradient: same shape as YYarray
Xˉ=∂L/∂X\bar X = \partial L / \partial Xthe VJP’s result: same shape as XXarray
dXdXa differential: an arbitrary small change of XXsame shape as XX
⟨A,B⟩=∑ijAijBij=tr⁡(A⊤B)\langle A, B\rangle = \sum_{ij} A_{ij} B_{ij} = \operatorname{tr}(A^\top B)the inner product of two same-shape arraysfloat
⊙\odotelementwise product
1\mathbf{1}the vector of onesfloat[D]
DDthe size of the normalized (last) axisint
μ,σ2\mu, \sigma^2mean and variance of a row over its DD entriesfloat per row
ϵ\epsilonsmall constant added to the variance (10−510^{-5} for LayerNorm, 10−610^{-6} for RMSNorm in this course’s models)float
rr (rstd)reciprocal standard deviation, 1/σ2+ϵ1/\sqrt{\sigma^2 + \epsilon} (LayerNorm) or 1/mean(x2)+ϵ1/\sqrt{\mathrm{mean}(x^2) + \epsilon} (RMSNorm)float per row
x^=(x−μ) r\hat x = (x - \mu)\, rthe normalized rowfloat[D]
γ,β,w\gamma, \beta, wlearned scale and shift (LayerNorm), learned scale (RMSNorm)float[D]
y=softmax(x)y = \mathrm{softmax}(x), ℓ=logsoftmax(x)\ell = \mathrm{logsoftmax}(x)probabilities and log-probabilities of a rowfloat[V]
tt, nna row’s target class; the number of rows whose target is not ignore_indexint

Gradients have the shape of their variable. Whatever layout convention a textbook uses, here Xˉ\bar X is stored with exactly the shape of XX, and the entry Xˉij\bar X_{ij} is ∂L/∂Xij\partial L / \partial X_{ij}. Then the first-order change of the loss is

dL=∑ijXˉij dXij=⟨Xˉ,dX⟩.dL = \sum_{ij} \bar X_{ij}\, dX_{ij} = \langle \bar X, dX\rangle.

The recipe. For Y=f(X)Y = f(X), the chain rule says dL=⟨G,dY⟩dL = \langle G, dY\rangle. Write dYdY as a linear expression in dXdX, then rearrange ⟨G,dY⟩\langle G, dY\rangle into the form ⟨something,dX⟩\langle \text{something}, dX\rangle. Since this holds for every dXdX, the something is Xˉ\bar X. Two tools do the rearranging:

  • the cyclic property of the trace: tr⁡(ABC)=tr⁡(CAB)=tr⁡(BCA)\operatorname{tr}(ABC) = \operatorname{tr}(CAB) = \operatorname{tr}(BCA), and tr⁡(A⊤)=tr⁡(A)\operatorname{tr}(A^\top) = \operatorname{tr}(A);
  • moving factors across the inner product: ⟨A,BC⟩=⟨B⊤A,C⟩=⟨AC⊤,B⟩\langle A, BC\rangle = \langle B^\top A, C\rangle = \langle A C^\top, B\rangle.

Matmul. Y=ABY = AB with A∈Rm×kA \in \mathbb{R}^{m \times k}, B∈Rk×nB \in \mathbb{R}^{k \times n}. The product rule holds for matrices (keep the order): dY=dA B+A dBdY = dA\, B + A\, dB. Then

⟨G,dA B⟩=tr⁡(G⊤dA B)=tr⁡(BG⊤dA)=⟨GB⊤,dA⟩,⟨G,A dB⟩=⟨A⊤G,dB⟩,\langle G, dA\, B\rangle = \operatorname{tr}(G^\top dA\, B) = \operatorname{tr}(B G^\top dA) = \langle G B^\top, dA\rangle, \qquad \langle G, A\, dB\rangle = \langle A^\top G, dB\rangle,

so Aˉ=GB⊤\bar A = G B^\top and Bˉ=A⊤G\bar B = A^\top G. Check the shapes: GG is m×nm \times n and B⊤B^\top is n×kn \times k, so Aˉ\bar A is m×km \times k like AA. Shape-checking catches many mistakes but not all: if BB is square, GBG B has the right shape and the wrong values.

Broadcasting. numpy’s matmul broadcasts batch axes: a weight BB of shape [k,n][k, n] times activations [T,m,k][T, m, k] is the same BB used TT times. Each use contributes a gradient, so Bˉ\bar B is the sum over the batch axis. In general, if the forward pass broadcast XX from shape ss to a larger shape, unbroadcast sums the gradient over every axis that was added in front and over every axis where ss has size 1 (keeping it as size 1).

Softmax. For one row, yi=exi/∑jexjy_i = e^{x_i} / \sum_j e^{x_j}. Differentiating the quotient gives dyi=yi (dxi−∑jyj dxj)dy_i = y_i\,(dx_i - \sum_j y_j\, dx_j), that is, dy=y⊙(dx−⟨y,dx⟩1)dy = y \odot (dx - \langle y, dx\rangle \mathbf{1}); the Jacobian is diag(y)−yy⊤\mathrm{diag}(y) - y y^\top. Then

⟨g,dy⟩=∑igiyi dxi−(∑igiyi)(∑jyj dxj)=⟨y⊙(g−⟨g,y⟩1),  dx⟩,\langle g, dy\rangle = \sum_i g_i y_i\, dx_i - \Big(\sum_i g_i y_i\Big)\Big(\sum_j y_j\, dx_j\Big) = \big\langle y \odot (g - \langle g, y\rangle \mathbf{1}),\; dx\big\rangle,

so xˉ=y⊙(g−⟨g,y⟩)\bar x = y \odot (g - \langle g, y\rangle). The VJP needs only the saved output yy; the diagonal term alone, y⊙gy \odot g, is a common half-derivation.

Log-softmax. ℓ=x−LSE(x)1\ell = x - \mathrm{LSE}(x)\mathbf{1} and d LSE=⟨y,dx⟩d\,\mathrm{LSE} = \langle y, dx\rangle (the gradient of log-sum-exp is softmax), so dℓ=dx−⟨y,dx⟩1d\ell = dx - \langle y, dx\rangle \mathbf{1} and

xˉ=g−y∑igi,y=eℓ.\bar x = g - y \sum_i g_i, \qquad y = e^{\ell}.

The function receives the saved log-probabilities ℓ\ell and exponentiates them.

Cross-entropy, fused. For logits zz (one row per position) and targets tt, the loss is the mean over the nn valid rows of −ℓt-\ell_{t}. For a valid row the upstream gradient of ℓ\ell is g=−1netg = -\tfrac1n e_t (a one-hot vector scaled), with ∑igi=−1n\sum_i g_i = -\tfrac1n. Plug into the log-softmax VJP:

zˉ=−1net+1ny=softmax(z)−onehot(t)n.\bar z = -\tfrac1n e_t + \tfrac1n y = \frac{\mathrm{softmax}(z) - \mathrm{onehot}(t)}{n}.

Rows whose target is ignore_index (padding, −100-100 as in PyTorch) are not in the loss, so their gradient is exactly 0 and they do not count in nn. Use M09.2’s softmax, so logits near 10410^4 give a finite gradient. A target outside [0,V)[0, V) that is not ignore_index is a data bug, not a class: numpy would quietly wrap −1-1 to the last column.

LayerNorm, defined. LayerNorm (Ba, Kiros, and Hinton, 2016) standardizes each row of features, then applies a learned scale and shift:

μ=1D∑ixi,σ2=1D∑i(xi−μ)2,r=1σ2+ϵ,x^=(x−μ) r,y=γ⊙x^+β.\mu = \tfrac1D \textstyle\sum_i x_i,\quad \sigma^2 = \tfrac1D \sum_i (x_i - \mu)^2,\quad r = \frac{1}{\sqrt{\sigma^2 + \epsilon}},\quad \hat x = (x - \mu)\,r,\quad y = \gamma \odot \hat x + \beta.

Every token’s features come out with mean 0 and variance about 1 before the scale, which keeps activations in a range where training is stable (the 2017 transformer of L5 uses it). The forward pass saves x^\hat x and rr.

LayerNorm’s VJP. The parameter gradients are immediate from y=γ⊙x^+βy = \gamma \odot \hat x + \beta: summing over every row, γˉ=∑rowsg⊙x^\bar\gamma = \sum_{\text{rows}} g \odot \hat x and βˉ=∑rowsg\bar\beta = \sum_{\text{rows}} g. For xx, let d=g⊙γd = g \odot \gamma be the gradient reaching x^\hat x. Both μ\mu and rr depend on every xix_i:

dμ=1D⟨1,dx⟩,dσ2=2D⟨x−μ,dx⟩  (because ∑i(xi−μ)=0),dr=−12r3 dσ2.d\mu = \tfrac1D \langle \mathbf{1}, dx\rangle, \qquad d\sigma^2 = \tfrac2D \langle x - \mu, dx\rangle \;(\text{because } \textstyle\sum_i (x_i - \mu) = 0), \qquad dr = -\tfrac12 r^3\, d\sigma^2.

So dx^=r (dx−dμ 1)+(x−μ) dr=r (dx−1D⟨1,dx⟩1)−rD x^ ⟨x^,dx⟩d\hat x = r\,(dx - d\mu\,\mathbf{1}) + (x - \mu)\,dr = r\,(dx - \tfrac1D\langle\mathbf{1}, dx\rangle\mathbf{1}) - \tfrac rD\, \hat x\, \langle \hat x, dx\rangle. Taking ⟨d,⋅⟩\langle d, \cdot\rangle and moving dxdx to the right:

xˉ=r(d−mean(d) 1−x^  mean(d⊙x^)).\bar x = r\left(d - \mathrm{mean}(d)\,\mathbf{1} - \hat x\; \mathrm{mean}(d \odot \hat x)\right).

Three terms: the direct path, the path through the mean, and the path through the variance. Dropping either correction is “treating μ\mu (or rr) as a constant”, and gradcheck catches it immediately. A consequence worth noticing: ⟨1,xˉ⟩=0\langle \mathbf{1}, \bar x\rangle = 0, because adding a constant to xx does not change x^\hat x.

RMSNorm, defined. RMSNorm (Zhang and Sennrich, 2019) drops the mean and the shift: it divides by the root mean square,

r=11D∑ixi2+ϵ,y=w⊙x r.r = \frac{1}{\sqrt{\tfrac1D \sum_i x_i^2 + \epsilon}}, \qquad y = w \odot x\, r.

It is cheaper and works as well in practice; Llama-family models (L7.1, SmolLM2) use it before every attention and MLP block.

RMSNorm’s VJP. With d=g⊙wd = g \odot w and dr=−12r3⋅2D⟨x,dx⟩=−r3D⟨x,dx⟩dr = -\tfrac12 r^3 \cdot \tfrac2D \langle x, dx\rangle = -\tfrac{r^3}{D} \langle x, dx\rangle:

⟨d,r dx+x dr⟩=r⟨d,dx⟩−r3D⟨d,x⟩⟨x,dx⟩  ⇒  xˉ=r(d−x r2 mean(d⊙x)),\langle d, r\,dx + x\,dr\rangle = r\langle d, dx\rangle - \tfrac{r^3}{D}\langle d, x\rangle \langle x, dx\rangle \;\Rightarrow\; \bar x = r\left(d - x\, r^2\, \mathrm{mean}(d \odot x)\right),

and wˉ=∑rowsg⊙x r\bar w = \sum_{\text{rows}} g \odot x\, r. Here ⟨x,xˉ⟩=0\langle x, \bar x\rangle = 0 when ϵ=0\epsilon = 0: scaling xx does not change the output.

Pure functions. A VJP reads the saved values and the upstream gradient and returns new arrays. Writing into its arguments (subtracting the one-hot from a softmax buffer you were handed, updating g in place) corrupts values the graph still needs.

Matmul. A=[1234]A = \begin{bmatrix}1&2\\3&4\end{bmatrix}, B=[1021]B = \begin{bmatrix}1&0\\2&1\end{bmatrix}, and G=[1000]G = \begin{bmatrix}1&0\\0&0\end{bmatrix}, so L=Y11=a11b11+a12b21L = Y_{11} = a_{11}b_{11} + a_{12}b_{21}. Directly: ∂L/∂a11=b11=1\partial L/\partial a_{11} = b_{11} = 1, ∂L/∂a12=b21=2\partial L/\partial a_{12} = b_{21} = 2, ∂L/∂b11=a11=1\partial L/\partial b_{11} = a_{11} = 1, ∂L/∂b21=a12=2\partial L/\partial b_{21} = a_{12} = 2, everything else 0. The formulas agree: GB⊤=[1200]G B^\top = \begin{bmatrix}1&2\\0&0\end{bmatrix} and A⊤G=[1020]A^\top G = \begin{bmatrix}1&0\\2&0\end{bmatrix}. Using GBGB instead gives [1000]\begin{bmatrix}1&0\\0&0\end{bmatrix}: the right shape, the wrong gradient.

Softmax. x=[0,ln⁡3]x = [0, \ln 3]: ex=[1,3]e^x = [1, 3], y=[1/4,3/4]y = [1/4, 3/4]. With g=[1,0]g = [1, 0], ⟨g,y⟩=1/4\langle g, y\rangle = 1/4, so xˉ=[14(1−14),34(0−14)]=[3/16,−3/16]\bar x = [\tfrac14(1 - \tfrac14), \tfrac34(0 - \tfrac14)] = [3/16, -3/16]. The entries sum to 0, as they must: adding a constant to xx does not change yy.

Cross-entropy. The same logits with target 1 and n=1n = 1: zˉ=y−e1=[1/4,−1/4]\bar z = y - e_1 = [1/4, -1/4]. The loss is −ln⁡(3/4)=0.2877-\ln(3/4) = 0.2877.

LayerNorm of x=[1,2,6]x = [1, 2, 6] with γ=1\gamma = \mathbf{1}, ϵ=0\epsilon = 0, g=[1,0,0]g = [1, 0, 0]:

quantityvalue
μ\mu, x−μx - \mu3, [−2,−1,3][-2, -1, 3]
σ2\sigma^2, rr14/314/3, 3/14=0.462910\sqrt{3/14} = 0.462910
x^\hat x[−0.925820,−0.462910,1.388730][-0.925820, -0.462910, 1.388730]
d=g⊙γd = g \odot \gamma, mean(d)\mathrm{mean}(d)[1,0,0][1, 0, 0], 1/31/3
mean(d⊙x^)\mathrm{mean}(d \odot \hat x)−0.925820/3=−0.308607-0.925820 / 3 = -0.308607
x^⋅mean(d⊙x^)\hat x \cdot \mathrm{mean}(d \odot \hat x)[2/7,1/7,−3/7][2/7, 1/7, -3/7]
$d - \mathrm{mean}(d) - $ that[8/21,−10/21,2/21][8/21, -10/21, 2/21]
xˉ=r⋅\bar x = r \cdot that[0.176347,−0.220433,0.044087][0.176347, -0.220433, 0.044087]

The entries of xˉ\bar x sum to 0. The parameter gradients are γˉ=g⊙x^=[−0.925820,0,0]\bar\gamma = g \odot \hat x = [-0.925820, 0, 0] and βˉ=[1,0,0]\bar\beta = [1, 0, 0].

RMSNorm of x=[3,4]x = [3, 4] with w=1w = \mathbf{1}, ϵ=0\epsilon = 0, g=[1,0]g = [1, 0]: mean(x2)=12.5\mathrm{mean}(x^2) = 12.5, r=0.282843r = 0.282843, y=[0.848528,1.131371]y = [0.848528, 1.131371]. Then d=[1,0]d = [1, 0], mean(d⊙x)=1.5\mathrm{mean}(d \odot x) = 1.5, r2=0.08r^2 = 0.08, xr2⋅1.5=[0.36,0.48]x r^2 \cdot 1.5 = [0.36, 0.48], and xˉ=r [0.64,−0.48]=[0.181019,−0.135765]\bar x = r\,[0.64, -0.48] = [0.181019, -0.135765]. Check: ⟨x,xˉ⟩∝3(0.64)+4(−0.48)=0\langle x, \bar x\rangle \propto 3(0.64) + 4(-0.48) = 0. And wˉ=g⊙x r=[0.848528,0]\bar w = g \odot x\,r = [0.848528, 0].

All of these are the first test in section 4, test_hand_example.

python/tinyllm/autograd/vjp.py
def unbroadcast(g, shape: tuple[int, ...]) -> NDArray
def matmul_vjp(g, A, B) -> tuple[NDArray, NDArray] # A [..., m, k], B [..., k, n]
def softmax_vjp(g, y, axis: int = -1) -> NDArray # y = softmax(x)
def log_softmax_vjp(g, y, axis: int = -1) -> NDArray # y = log_softmax(x)
def layernorm_vjp(g, xhat, rstd, gamma) -> tuple[NDArray, NDArray, NDArray] # dx, dgamma, dbeta
def rmsnorm_vjp(g, x, rstd, w) -> tuple[NDArray, NDArray] # dx, dw
def cross_entropy_vjp(logits, targets, ignore_index: int = -100) -> NDArray

LayerNorm and RMSNorm normalize the last axis; rstd may be passed with shape x.shape[:-1] or x.shape[:-1] + (1,), and the parameter gradients sum over every leading axis. matmul_vjp needs at least 2-D operands and unbroadcasts both gradients. The forward passes are not part of this module: L0.2 writes them and saves y, xhat, and rstd; the tests here compute their own.

TestKINDChecksWhy it matters downstream
test_hand_exampleunitevery section 3 numberyou and the test agree on each rule
test_matches_torch_goldengoldentorch.autograd’s gradients for all seven rules, in float64the framework L0.2 is compared against
test_unbroadcast_shapesunitseven broadcast patterns, counted with a gradient of onesbiases and shared weights
test_unbroadcast_rejects_impossible_shapesboundaryshapes that could not have broadcast raise ValueErrora wrong-shaped gradient is a bug
test_matmul_vjp_gradcheckgradcheck2-D, square, and batched products with either operand broadcastevery linear layer
test_matmul_vjp_rejects_vectorsboundary1-D operands raise ValueErrorvectors must be reshaped explicitly
test_softmax_vjps_gradcheckgradchecksoftmax and log-softmax along the last axis and axis 0attention weights, the loss
test_softmax_vjp_matches_numeric_vjpdifferentialagainst M04.2’s vjp_numeric on one rowthe same u⊤Ju^\top J two ways
test_layernorm_vjp_gradcheckgradcheckxx, γ\gamma, β\beta on [2, 3, 8], rstd in both shapesthe 2017 transformer (L5)
test_rmsnorm_vjp_gradcheckgradcheckxx and ww, rstd in both shapesL7.1
test_cross_entropy_vjp_gradcheckgradcheckthe gradient of M11.1’s cross-entropy, with ignored rowsL0.3’s fused loss
test_cross_entropy_ignore_indexboundaryignored rows get 0, the mean counts valid rows, out-of-range targets raisepadded batches
test_cross_entropy_large_logitsboundarylogits near 10410^4 give a finite gradienta confident model
test_trace_identitypropertytr⁡(G⊤dY)=⟨Aˉ,dA⟩+⟨Bˉ,dB⟩\operatorname{tr}(G^\top dY) = \langle \bar A, dA\rangle + \langle \bar B, dB\rangle on 20 directionsthe definition of a VJP
test_vjps_do_not_mutate_inputsunitsaved values and upstream gradients are unchangedthe graph reuses them
PitfallSymptomCaught by
1. GBG B instead of GB⊤G B^\topcorrect shape for square BB, wrong valuestest_matmul_vjp_gradcheck (mutant s01)
2. a transposed gradientBˉ⊤\bar B^\top where Bˉ\bar B belongstest_trace_identity (mutant s02)
3. not summing over broadcast axesa bias or shared weight gets a batch of gradientstest_matmul_vjp_gradcheck (mutant s03), test_unbroadcast_shapes (mutants s14, s15)
4. the softmax Jacobian’s diagonal onlyxˉ=y⊙g\bar x = y \odot gtest_softmax_vjps_gradcheck (mutant s04)
5. log-softmax VJP with ℓ\ell where eℓe^\ell belongslog-probabilities used as probabilitiestest_softmax_vjps_gradcheck (mutant s05)
6. treating μ\mu or rr as a constantmissing correction terms in LayerNorm or RMSNormtest_layernorm_vjp_gradcheck (mutants s06, s07), test_rmsnorm_vjp_gradcheck (mutants s09, s10)
7. γˉ=∑g\bar\gamma = \sum gthe scale learns like a shifttest_layernorm_vjp_gradcheck (mutant s08)
8. averaging over every row, or letting padding throughgradients scaled by the padding ratio; padding tokens trainedtest_cross_entropy_ignore_index (mutants s11, s12, s13)
9. exp(z) / sum(exp(z)) inside the loss gradientNaN once a logit passes 709test_cross_entropy_large_logits (mutant s16)
10. updating the upstream gradient in placethe caller’s g changes under ittest_vjps_do_not_mutate_inputs (mutant s17)
DirectionModuleHow it uses this
BackM09.2softmax inside cross_entropy_vjp
BackM11.1its cross_entropy is the loss whose gradient cross_entropy_vjp is
BackM04.2vjp_numeric checks the softmax VJP numerically
BackM03.1row-major matrices and matmul shapes
ForwardL0.1your Tensor’s backward sums broadcast gradients back to each input’s shape with unbroadcast
ForwardL0.2each op of the tensor library registers one of these VJPs
ForwardL0.3the fused softmax cross-entropy, (softmax−onehot)/n(\mathrm{softmax} - \mathrm{onehot})/n
ForwardL3.1backpropagation through time by hand chains matmul_vjp over steps
ForwardL7.1rmsnorm_vjp is RMSNorm’s backward

If you skip this module, ol check L0.2 stops with L0.2 needs M08.3: build it, or rerun with --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
the closed-form rulesPyTorch’s derivative tableone backward formula per op, from which the autograd code is generatedtools/autograd/derivatives.yaml
layernorm_vjpllm.cthe same three-term backward in plain C, fused over a batchtrain_gpt2.c (layernorm_backward)
cross_entropy_vjpLiger Kernelthe linear layer, softmax, and cross-entropy fused in chunks, so the [T,V][T, V] logits never exist in memorysrc/liger_kernel/ops/fused_linear_cross_entropy.py
rmsnorm_vjpPyTorch F.rms_normfused kernels with the saved rstd, the same formulaaten/src/ATen/native/layer_norm.cpp