Matrix differentials, the trace trick, and closed-form VJPs
Overview
Section titled “Overview”| Module | M08.3 · build · Python · Pass 2 · 3 to 4 h |
| You build | python/tinyllm/autograd/vjp.py: unbroadcast, matmul_vjp, softmax_vjp, log_softmax_vjp, layernorm_vjp, rmsnorm_vjp, cross_entropy_vjp |
| Contract | course/contracts/py/tinyllm/autograd/vjp.pyi |
| Tests | course/tests/M08.3/ (what they check: section 4) |
| Needs | M09.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 by | L0.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 |
| Milestone | MS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains) |
| Optional depth | Parr 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 |
Key Takeaways
Section titled “Key Takeaways”- For a scalar loss, with ; write in terms of , move to the right with the cyclic property of the trace, and the matrix in front of it is (
test_trace_identity). - That gives every rule in a few lines: and for a matmul, for softmax, 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).
How to work this chapter
Section titled “How to work this chapter”ol start M08.3 # stubs vjp.py into your repo, contract alongsideol tests M08.3 # read the test catalog first: rung R0, you write no tests hereol check M08.3 # exit code is the verdictol check M08.3 --ref-deps # only if your M09.2, M11.1, or M04.2 is not passing yetol diff M08.3 # after passing: your code against the reference1. Why now
Section titled “1. Why now”M08.2 gave you reverse mode one scalar at a time. Your bigram’s forward pass is a 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.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| a scalar loss | float | |
| an op’s input and output | arrays | |
| the upstream gradient: same shape as | array | |
| the VJP’s result: same shape as | array | |
| a differential: an arbitrary small change of | same shape as | |
| the inner product of two same-shape arrays | float | |
| elementwise product | ||
| the vector of ones | float[D] | |
| the size of the normalized (last) axis | int | |
| mean and variance of a row over its entries | float per row | |
| small constant added to the variance ( for LayerNorm, for RMSNorm in this course’s models) | float | |
(rstd) | reciprocal standard deviation, (LayerNorm) or (RMSNorm) | float per row |
| the normalized row | float[D] | |
| learned scale and shift (LayerNorm), learned scale (RMSNorm) | float[D] | |
| , | probabilities and log-probabilities of a row | float[V] |
| , | a row’s target class; the number of rows whose target is not ignore_index | int |
Gradients have the shape of their variable. Whatever layout convention a textbook uses, here is stored with exactly the shape of , and the entry is . Then the first-order change of the loss is
The recipe. For , the chain rule says . Write as a linear expression in , then rearrange into the form . Since this holds for every , the something is . Two tools do the rearranging:
- the cyclic property of the trace: , and ;
- moving factors across the inner product: .
Matmul. with , . The product rule holds for matrices (keep the order): . Then
so and . Check the shapes: is and is , so is like . Shape-checking catches many mistakes but not all: if is square, has the right shape and the wrong values.
Broadcasting. numpy’s matmul broadcasts batch axes: a weight of shape times activations is the same used times. Each use contributes a gradient, so is the sum over the batch axis. In general, if the forward pass broadcast from shape to a larger shape, unbroadcast sums the gradient over every axis that was added in front and over every axis where has size 1 (keeping it as size 1).
Softmax. For one row, . Differentiating the quotient gives , that is, ; the Jacobian is . Then
so . The VJP needs only the saved output ; the diagonal term alone, , is a common half-derivation.
Log-softmax. and (the gradient of log-sum-exp is softmax), so and
The function receives the saved log-probabilities and exponentiates them.
Cross-entropy, fused. For logits (one row per position) and targets , the loss is the mean over the valid rows of . For a valid row the upstream gradient of is (a one-hot vector scaled), with . Plug into the log-softmax VJP:
Rows whose target is ignore_index (padding, as in PyTorch) are not in the loss, so their gradient is exactly 0 and they do not count in . Use M09.2’s softmax, so logits near give a finite gradient. A target outside that is not ignore_index is a data bug, not a class: numpy would quietly wrap to the last column.
LayerNorm, defined. LayerNorm (Ba, Kiros, and Hinton, 2016) standardizes each row of features, then applies a learned scale and shift:
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 and .
LayerNorm’s VJP. The parameter gradients are immediate from : summing over every row, and . For , let be the gradient reaching . Both and depend on every :
So . Taking and moving to the right:
Three terms: the direct path, the path through the mean, and the path through the variance. Dropping either correction is “treating (or ) as a constant”, and gradcheck catches it immediately. A consequence worth noticing: , because adding a constant to does not change .
RMSNorm, defined. RMSNorm (Zhang and Sennrich, 2019) drops the mean and the shift: it divides by the root mean square,
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 and :
and . Here when : scaling 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.
3. Worked example by hand
Section titled “3. Worked example by hand”Matmul. , , and , so . Directly: , , , , everything else 0. The formulas agree: and . Using instead gives : the right shape, the wrong gradient.
Softmax. : , . With , , so . The entries sum to 0, as they must: adding a constant to does not change .
Cross-entropy. The same logits with target 1 and : . The loss is .
LayerNorm of with , , :
| quantity | value |
|---|---|
| , | 3, |
| , | , |
| , | , |
| $d - \mathrm{mean}(d) - $ that | |
| that |
The entries of sum to 0. The parameter gradients are and .
RMSNorm of with , , : , , . Then , , , , and . Check: . And .
All of these are the first test in section 4, test_hand_example.
4. The interface
Section titled “4. The interface”def unbroadcast(g, shape: tuple[int, ...]) -> NDArraydef 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, dbetadef rmsnorm_vjp(g, x, rstd, w) -> tuple[NDArray, NDArray] # dx, dwdef cross_entropy_vjp(logits, targets, ignore_index: int = -100) -> NDArrayLayerNorm 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.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example | unit | every section 3 number | you and the test agree on each rule |
test_matches_torch_golden | golden | torch.autograd’s gradients for all seven rules, in float64 | the framework L0.2 is compared against |
test_unbroadcast_shapes | unit | seven broadcast patterns, counted with a gradient of ones | biases and shared weights |
test_unbroadcast_rejects_impossible_shapes | boundary | shapes that could not have broadcast raise ValueError | a wrong-shaped gradient is a bug |
test_matmul_vjp_gradcheck | gradcheck | 2-D, square, and batched products with either operand broadcast | every linear layer |
test_matmul_vjp_rejects_vectors | boundary | 1-D operands raise ValueError | vectors must be reshaped explicitly |
test_softmax_vjps_gradcheck | gradcheck | softmax and log-softmax along the last axis and axis 0 | attention weights, the loss |
test_softmax_vjp_matches_numeric_vjp | differential | against M04.2’s vjp_numeric on one row | the same two ways |
test_layernorm_vjp_gradcheck | gradcheck | , , on [2, 3, 8], rstd in both shapes | the 2017 transformer (L5) |
test_rmsnorm_vjp_gradcheck | gradcheck | and , rstd in both shapes | L7.1 |
test_cross_entropy_vjp_gradcheck | gradcheck | the gradient of M11.1’s cross-entropy, with ignored rows | L0.3’s fused loss |
test_cross_entropy_ignore_index | boundary | ignored rows get 0, the mean counts valid rows, out-of-range targets raise | padded batches |
test_cross_entropy_large_logits | boundary | logits near give a finite gradient | a confident model |
test_trace_identity | property | on 20 directions | the definition of a VJP |
test_vjps_do_not_mutate_inputs | unit | saved values and upstream gradients are unchanged | the graph reuses them |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. instead of | correct shape for square , wrong values | test_matmul_vjp_gradcheck (mutant s01) |
| 2. a transposed gradient | where belongs | test_trace_identity (mutant s02) |
| 3. not summing over broadcast axes | a bias or shared weight gets a batch of gradients | test_matmul_vjp_gradcheck (mutant s03), test_unbroadcast_shapes (mutants s14, s15) |
| 4. the softmax Jacobian’s diagonal only | test_softmax_vjps_gradcheck (mutant s04) | |
| 5. log-softmax VJP with where belongs | log-probabilities used as probabilities | test_softmax_vjps_gradcheck (mutant s05) |
| 6. treating or as a constant | missing correction terms in LayerNorm or RMSNorm | test_layernorm_vjp_gradcheck (mutants s06, s07), test_rmsnorm_vjp_gradcheck (mutants s09, s10) |
| 7. | the scale learns like a shift | test_layernorm_vjp_gradcheck (mutant s08) |
| 8. averaging over every row, or letting padding through | gradients scaled by the padding ratio; padding tokens trained | test_cross_entropy_ignore_index (mutants s11, s12, s13) |
9. exp(z) / sum(exp(z)) inside the loss gradient | NaN once a logit passes 709 | test_cross_entropy_large_logits (mutant s16) |
| 10. updating the upstream gradient in place | the caller’s g changes under it | test_vjps_do_not_mutate_inputs (mutant s17) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | M09.2 | softmax inside cross_entropy_vjp |
| Back | M11.1 | its cross_entropy is the loss whose gradient cross_entropy_vjp is |
| Back | M04.2 | vjp_numeric checks the softmax VJP numerically |
| Back | M03.1 | row-major matrices and matmul shapes |
| Forward | L0.1 | your Tensor’s backward sums broadcast gradients back to each input’s shape with unbroadcast |
| Forward | L0.2 | each op of the tensor library registers one of these VJPs |
| Forward | L0.3 | the fused softmax cross-entropy, |
| Forward | L3.1 | backpropagation through time by hand chains matmul_vjp over steps |
| Forward | L7.1 | rmsnorm_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.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
| the closed-form rules | PyTorch’s derivative table | one backward formula per op, from which the autograd code is generated | tools/autograd/derivatives.yaml |
layernorm_vjp | llm.c | the same three-term backward in plain C, fused over a batch | train_gpt2.c (layernorm_backward) |
cross_entropy_vjp | Liger Kernel | the linear layer, softmax, and cross-entropy fused in chunks, so the logits never exist in memory | src/liger_kernel/ops/fused_linear_cross_entropy.py |
rmsnorm_vjp | PyTorch F.rms_norm | fused kernels with the saved rstd, the same formula | aten/src/ATen/native/layer_norm.cpp |