Matrix calculus problem set: differentials, trace trick, VJPs, SDPA Jacobian
Overview
Section titled “Overview”| Module | S-M08 · solve · none · Pass 2 · 8 to 10 h |
| You build | answers in solve/S-M08.toml (42 checked by SymPy) and 4 derivations in solve/S-M08/qN.md (self-graded against their rubrics) |
| Contract | none: a pen and paper set |
| Tests | course/solve/S-M08/key.toml (hidden): typed answers plus reject canaries; the problems are in course/solve/S-M08/problems.md and in section 4 |
| Needs | S-M05 (proof habits). Reading: the Matrix Calculus and Autodiff topic, and the gradients and chain rule of S-M04 |
| Used by | no call site (a solve set). It checks the derivations behind M08.3 (matmul_vjp, softmax_vjp, log_softmax_vjp, layernorm_vjp, rmsnorm_vjp, cross_entropy_vjp), which L0.2, L0.3, and L7.1 register as ops; q33 and q34 are the attention backward of L4.* and L9.4 |
| Milestone | MS-P2 (the Pass 2 gate runs ol check on every solve part of the pass) |
| Optional depth | Parr and Howard, “The Matrix Calculus You Need for Deep Learning” (2018); Minka, “Old and New Matrix Algebra Useful for Statistics” (2000); Baydin et al., “Automatic Differentiation in Machine Learning: a Survey” (2018), sections 2 and 3; Dao et al., “FlashAttention” (2022), appendix B |
Key Takeaways
Section titled “Key Takeaways”- A gradient has the shape of its variable, and broadcasting in the forward pass becomes a sum over the broadcast axes in the backward pass (q1, q3).
- Write the differential , push through the product rule, rotate with the cyclic trace, and read off the gradient: , (q18, q20).
- Softmax’s VJP is and cross-entropy’s gradient with respect to logits is , both (q21, q24, q28).
- Normalizations differentiate through their statistics: LayerNorm’s input gradient sums to zero, and RMSNorm’s is orthogonal to (q26, q27).
- Reverse mode costs one pass per output, so one backward pass gives a million-parameter gradient, at about twice the forward FLOPs (q29, q32).
How to work this chapter
Section titled “How to work this chapter”ol start S-M08 # writes solve/S-M08.toml and one file per derivationol check S-M08 # SymPy checks the answers, then asks each rubric (y/n)ol check S-M08 --regrade # ask the rubrics again after you change a derivation1. Why now
Section titled “1. Why now”By this point in Pass 2 you have a scalar autograd engine (M08.2) that differentiates one number at a time. It is correct and hopeless at scale: applying one weight to one vector would create a third of a million scalar multiplication nodes. M08.3 replaces it with closed-form vector-Jacobian products, one per operation, that take the upstream gradient as an array and return the input gradients as arrays, and L0.2 registers them as the ops of your tensor autograd. Each closed form is a derivation you must get right on paper first, because a gradient with the right shape and the wrong value trains, slowly and badly, rather than crashing. This set drills the shape bookkeeping, matrix differentials, the trace trick, the VJPs of softmax, log-softmax, cross-entropy, LayerNorm, and RMSNorm, the cost model of forward and reverse mode, and the attention backward you will meet again in L9.4.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| a scalar loss | scalar | |
| , | matrices; | |
| upstream gradient, same shape as | ||
| differential: first-order change of | like | |
| trace, | scalar | |
| Frobenius inner product, | scalar | |
| elementwise product | ||
| Jacobian of , | ||
| vector | ||
| 1 if , else 0 | ||
| the all-ones vector |
2.1 Shapes, layouts, and strides
Section titled “2.1 Shapes, layouts, and strides”The convention in this course (and in PyTorch) is the denominator layout: has the shape of , entry by entry . A row-major array of shape stores index at element offset with strides . A view reinterprets the same memory with other strides: swapping two axes swaps their strides and copies nothing. Broadcasting reuses one value along an axis; its backward is a sum over that axis, which L0.1’s unbroadcast performs. A Jacobian of a map from to is , which for a layer is enormous; reverse mode never builds it.
2.2 Differentials
Section titled “2.2 Differentials”The differential of at is the linear part of . It obeys the rules you know for scalars, with the order of factors kept: , , , , and (q16). For the result is . For an elementwise function , , so its Jacobian is diagonal.
2.3 The trace trick
Section titled “2.3 The trace trick”For a scalar of a matrix , . So once is written as , the gradient is . The tool for getting there is the cyclic property: whenever the products are defined (a cyclic shift, never an arbitrary swap), together with . A scalar is its own trace, so any scalar expression may be wrapped in one.
2.4 VJPs of the ops you will register
Section titled “2.4 VJPs of the ops you will register”A vector-Jacobian product maps an upstream gradient to without forming . The ones M08.3 implements:
| Op | Forward | VJP |
|---|---|---|
| matmul | , | |
| softmax | ||
| log-softmax | ||
| cross-entropy | ; a mean over kept positions divides by | |
| RMSNorm | , | |
| LayerNorm | , | with : |
Here is the mean of the entries of , the mean of , , and the one-hot vector of the target. LayerNorm and RMSNorm are functions of a whole row: LayerNorm centers and scales it to mean 0 and variance 1, RMSNorm only scales it to root-mean-square 1, and both then apply a learned gain. Because each normalizes by a statistic of its own input, the gradient has a correction term that a “treat the statistic as constant” derivation misses. Two checks catch it: LayerNorm’s input gradient sums to zero (shifting changes nothing), and RMSNorm’s is orthogonal to when (scaling changes nothing).
2.5 Forward versus reverse mode
Section titled “2.5 Forward versus reverse mode”For , forward mode (M08.1’s dual numbers) computes one Jacobian-vector product per pass, so the full Jacobian takes passes. Reverse mode computes one per pass, so it takes . A loss has : one backward pass gives the gradient with respect to every parameter, at the price of storing the forward activations it needs. For a matrix product the backward computes two products of the same size as the forward one, so backward costs about twice the forward FLOPs, and a training step about three times: the per token of S-M05 q13.
2.6 Attention
Section titled “2.6 Attention”Scaled dot-product attention (SDPA) for positions computes scores , weights along each row, and output . Its backward chains the three rules above: matmul for , softmax row by row, matmul again for (q34).
3. Worked example by hand
Section titled “3. Worked example by hand”This is a sibling of q21 and q24, not one of the graded problems.
Softmax VJP. Let , so , and let the upstream gradient be .
- .
- .
- .
Check against the Jacobian: for two classes , and . The entries sum to 0, as every softmax input gradient must. In solve/ this is answer = "[-1/2, 1/2]".
Cross-entropy. Logits , target class 0. Then , , and : lower the wrong logit, raise the right one, by how far the prediction is from the target.
4. The problem set
Section titled “4. The problem set”Write each answer in solve/S-M08.toml; lettered parts are their own tables:
[q1.a]answer = "[2, 3, 5]"[q5]answer = "[2*x1 + 2*x2, 2*x1 + 6*x2]"[q18.a]answer = "[[1, 3], [2, 4]]"[q20]proof = "S-M08/q20.md"Vectors are flat lists; matrices are lists of rows; write as E and as sqrt(...). A gradient answer must have the variable’s shape.
Layouts and shapes
Section titled “Layouts and shapes”A linear layer computes for a batch of shape , stored row-major in float32, with of shape and of shape . is a scalar loss.
q1. Give the shape of (a) , (b) , (c) . [vector]
q2. (a) At which element offset (counting from 0) is stored? [number] (b) Give the strides of in elements, one per axis. [vector] (c) Give the strides of the view that swaps the last two axes of (shape , no copy). [vector]
q3. is broadcast over the batch and time axes. How many entries of are summed into each entry of ? [number]
q4. Flatten and row-major into vectors. What is the shape of the Jacobian (rows, columns)? [vector]
Differentials of matrix expressions
Section titled “Differentials of matrix expressions”q5. with and . Give . [vector in x1, x2]
q6. with . Give . [matrix]
q7. with , . Give at . [vector]
q8. . Give at . [matrix]
q9. Give at . [matrix]
q10. with . Give at . [matrix]
q11. Give the derivative of the sigmoid . [expr in z]
q12. Give the derivative of softplus, . [expr in z]
q13. Give the gradient of . [vector in x1, x2]
q14. Give the gradient of at . [vector]
q15. Give the Jacobian of elementwise at . [matrix]
q16. Prove for an invertible square matrix . [proof]
The trace trick
Section titled “The trace trick”q17. For conformable matrices: (a) is always? (b) Is always? [bool]
q18. with , , and upstream gradient . Give (a) and (b) . [matrix]
q19. The Frobenius inner product is . Give it for , . [number]
q20. Let and a scalar with , . Prove and . [proof]
VJPs by hand
Section titled “VJPs by hand”q21. and the upstream gradient is . Give . [vector]
q22. Give the Jacobian of a 2-class softmax at . [matrix]
q23. with and upstream gradient . Give . [vector]
q24. Cross-entropy with logits and target class 1 (counting from 0): . (a) Give . [vector] (b) Give . [number]
q25. A batch has 4 positions; one has the target ignore_index and the loss is the mean over the other positions. By what factor is each kept position’s per-position gradient scaled? [number]
q26. RMSNorm with : , . (a) For , , and upstream , give . [vector] (b) Is for every and every ? [bool]
q27. LayerNorm with , , : with the mean and . For and upstream , give . [vector]
q28. Prove the softmax VJP: if and , then . [proof]
Forward versus reverse mode
Section titled “Forward versus reverse mode”q29. A loss has parameters. How many (a) Jacobian-vector products (forward mode) and (b) vector-Jacobian products (reverse mode) give the full gradient? [number]
q30. You need the full Jacobian of . Which mode needs fewer passes? (a) forward (b) reverse [choice]
q31. Backward of a 16-layer MLP stores each layer’s input: a batch of 32 rows of width 1024 in float32. How many bytes of activations does it keep? [number]
q32. with of shape and of shape . (a) How many FLOPs does backward take to compute both and ? [expr in m, k, n] (b) What is the ratio (forward + backward) / forward? [number]
The SDPA Jacobian
Section titled “The SDPA Jacobian”Scaled dot-product attention: , row by row, . Take positions, , , , , no mask. Write as E.
q33. (a) Give . [matrix] With upstream gradient , give (b) and (c) . [matrix]
q34. Prove the SDPA backward rules: with , , , , , and . [proof]
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| Gradient transposed relative to its variable | shape error in the optimizer, or a silently wrong square gradient | q1 (canary [4, 5]), q9 (canary: the inverse) |
| Not summing over broadcast axes | bias gradient of shape (B, T, d) | q1, q3 (canaries 30 and 2) |
| Counting bytes as strides, or using a contiguous copy’s strides for a view | wrong element read through a view | q2 (canaries in bytes and [12, 3, 1]) |
| for a non-symmetric | wrong gradient for asymmetric forms | q5 (canary 2Ax) |
| Dropping the factor 2 of a square | gradients half as large; learning rates tuned to hide it | q7, q10 (canaries) |
| Swapping a non-cyclic order inside a trace | an identity that holds only for commuting matrices | q17 (canary true) |
| Softmax VJP without the term | input gradients that do not sum to zero | q21 (canary y * g) |
| One-hot minus softmax | gradient ascent on the loss | q24 (canary with the sign flipped) |
| Mean over all positions including ignored ones | loss and gradient scaled down by the padding fraction | q25 (canary 1/4) |
| Treating a normalization’s statistics as constants | RMSNorm or LayerNorm gradients without the correction term | q26 and q27 (canaries) |
| Using for | value gradients from the wrong positions | q33 (canary) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | S-M05 | proof structure for q16, q20, q28, q34 |
| Forward | M08.1 | dual numbers compute the JVPs of section 2.5 |
| Forward | M08.2 | scalar reverse mode, the oracle against which M08.3 is tested |
| Forward | M08.3 | the VJP table of section 2.4 as code, gradchecked rule by rule |
| Forward | L0.2 | registers those VJPs as tensor ops, with unbroadcast from q1 and q3 |
| Forward | L0.3 | fused cross-entropy with ignore_index (q24, q25) |
| Forward | L7.1 | RMSNorm in the modern block (q26) |
| Forward | L9.4 | FlashAttention backward recomputes and applies q34’s rules tile by tile |