Skip to content

Matrix Calculus and Autodiff

  • Autodiff is the chain rule applied to a program, not to a formula. Forward mode carries a derivative alongside each value; reverse mode records the computation and walks it backward once.
  • Reverse mode costs about one extra forward pass per scalar output, which is why every neural network trains with it: one loss, millions of parameters.
  • Every op needs one rule: its vector-Jacobian product (VJP). The trace trick and matrix differentials derive the VJPs of matmul, softmax, LayerNorm, RMSNorm, and cross-entropy in a few lines each.
  • Gradients are checked, never trusted. Central differences in float64 against the analytic VJP catch nearly every wrong rule.
  • Memory is the price of reverse mode. Activation checkpointing trades recomputation for memory on a n\sqrt{n} schedule.

Read Parr and Howard first, deriving every identity by hand, then the Baydin survey sections 2 and 3. Build micrograd from memory once before the course modules. In the course, S-M08 checks the derivations and M08.1 to M08.3 are the code: dual numbers, a scalar Value, then the closed-form VJPs that L0.2 registers as ops.


A gradient is a linear map, and a program is a composition of small differentiable steps. Forward mode pushes a tangent through each step (a Jacobian-vector product); reverse mode pulls a cotangent back through each step (a vector-Jacobian product). Knowing the VJP of each primitive is enough to differentiate any program built from them, and knowing it in closed form is what makes it fast and stable.

Key ideas:

  • Dual number a+bεa + b\varepsilon with ε2=0\varepsilon^2 = 0: evaluating f(x+ε)=f(x)+f′(x)εf(x + \varepsilon) = f(x) + f'(x)\varepsilon gives value and derivative together.
  • Cost: one pass per input direction, so it suits few inputs and many outputs.
  • Call site: M08.1 is the derivative oracle for the elementwise ops of L0.2.

Key ideas:

  • Tape: record each op and its inputs; backward visits nodes in reverse topological order (M06.1) and accumulates xˉ+=yˉ ∂y/∂x\bar{x} \mathrel{+}= \bar{y}\, \partial y / \partial x.
  • Broadcasting: the gradient of a broadcast input is the sum of the upstream gradient over the broadcast axes (unbroadcast).
  • Call site: M08.2 is the scalarized oracle for L0.1 broadcasting backward.

3. Matrix differentials and closed-form VJPs

Section titled “3. Matrix differentials and closed-form VJPs”

Key ideas:

  • Trace trick: for scalar LL, dL=tr⁡(G⊤dY)dL = \operatorname{tr}(G^{\top} dY); with Y=ABY = AB, dY=dA B+A dBdY = dA\,B + A\,dB gives Aˉ=GB⊤\bar{A} = G B^{\top}, Bˉ=A⊤G\bar{B} = A^{\top} G.
  • Softmax: xˉ=y⊙(g−⟨g,y⟩)\bar{x} = y \odot (g - \langle g, y \rangle); fused with cross-entropy the gradient is softmax(z)−onehot(t)\text{softmax}(z) - \text{onehot}(t).
  • Normalization: LayerNorm and RMSNorm VJPs reuse the saved reciprocal standard deviation; beat 2 of M08.3 defines both functions before deriving them.

Key ideas:

  • Hessian-vector product by finite differences of gradients, or forward-over-reverse, without forming the Hessian.
  • Checkpointing: keep activations at segment boundaries and recompute inside a segment; n\sqrt{n} segments balance memory and time (M08.4, used by L11.1).
ModuleTopicKindPass
S-M08Matrix calculus problem set (VJP derivations, checked by SymPy)solve2
M08.1Dual numbers, forward modebuild2
M08.2Scalar reverse mode (Value)build2
M08.3Matrix differentials, trace trick, closed-form VJPs. Beat 2 defines LayerNorm and RMSNorm as functions before deriving their VJPsbuild2
M08.4Hessian-vector products, recompute vs memory schedulebuild9
#ModuleChapterKindPass
1M08.1Dual numbers and forward-mode autodiffbuild2
2M08.2Scalar reverse mode: the Value graphbuild2
3M08.3Matrix differentials, the trace trick, and closed-form VJPsbuild2
4M08.4Hessian-vector products and the recompute-versus-memory schedulebuild9
5S-M08Matrix calculus problem set: differentials, trace trick, VJPs, SDPA Jacobiansolve2
TrackConnection
Calculus 3gradients, Jacobians, and the multivariable chain rule; M04.1 gradcheck
Numerical Methods and Floating Pointstable softmax and log-sum-exp inside the VJPs
Optimizationthe optimizers consume these gradients
tinyllm Part 0L0.1 to L0.3 turn these rules into the learner’s autograd engine
Deep Learningbackpropagation in the textbook setting