Hessian-vector products and the recompute-versus-memory schedule
Overview
Section titled “Overview”| Module | M08.4 · build · Python · Pass 9 · 3 h |
| You build | python/tinyllm/autograd/hvp.py: hvp_fd, hessian_fd, checkpoint_cost, min_checkpoint_memory, checkpoint_schedule |
| Contract | course/contracts/py/tinyllm/autograd/hvp.pyi |
| Tests | course/tests/M08.4/ (what they check: section 4) |
| Needs | nothing to build first · reading: M08.3 closed-form VJPs (the golden tests differentiate its cross-entropy gradient), S-M08 forward versus reverse cost, M04.2 the Hessian and the second-derivative test |
| Used by | L11.1 activation checkpointing plans its segments with checkpoint_schedule · later: M10.5 (optional) estimates the top Hessian eigenvalue with hvp_fd |
| Milestone | MS-P9 (Pass 9 gate: the math of the pass checks green, then the capstone trains with checkpointing) |
| Optional depth | Pearlmutter, “Fast Exact Multiplication by the Hessian” (Neural Computation, 1994); Chen, Xu, Zhang, and Guestrin, Training Deep Nets with Sublinear Memory Cost (2016); Griewank and Walther, Evaluating Derivatives, chapter 12 (checkpointing) |
Key Takeaways
Section titled “Key Takeaways”- The Hessian-vector product is the derivative of the gradient along , so two gradient calls give it to without forming the Hessian (
test_hand_example_quartic,test_error_shrinks_as_eps_squared,test_matches_torch_hvp). - On a quadratic the central difference is exact at every , because the gradient is affine (
test_quadratic_hvp_is_exact). - Reverse mode keeps one saved input per layer. Recomputing segments trades forward time for memory: the peak is and the extra work is every layer outside the last segment (
test_checkpoint_cost_extremes). - Uniform segments of layers peak near ; shrinking segments (sizes ) reach the optimum , and under a budget the best schedule is found greedily (
test_min_memory_is_triangular,test_schedule_is_optimal).
How to work this chapter
Section titled “How to work this chapter”ol start M08.4 # stubs hvp.py into your repo, contract alongsideol tests M08.4 # read the test catalog first: rung R0, you write no tests hereol check M08.4 # exit code is the verdictol diff M08.4 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Pass 9 trains the capstone, an 8-layer Llama with a 512-token context, on a laptop. Its backward pass needs every layer’s saved activations at once: at batch 16 that is several hundred megabytes in float32 before the logits, and the activations, not the 10M weights, are what runs the process out of memory. L11.1 fixes it with activation checkpointing: keep only some layer inputs, recompute the rest during backward. Which layers to keep is a counting problem with an exact answer, and you need that answer before you write the mechanism. The same pass is where curvature questions come up (“is my learning rate past ?”, M10.5), and the tool for those is the Hessian-vector product, which never forms the Hessian. This module gives you both: hvp_fd and the schedule functions L11.1 calls.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| a scalar loss of parameters | function | |
| the point (parameters) | float64[...] | |
the gradient (grad_fn) | same shape as | |
| the Hessian: | , symmetric | |
| a direction | same shape as | |
the finite-difference step along (eps) | float | |
| the -th directional derivative of along (a derivative of order of ) | same shape as | |
| the -th unit vector | float64[n] | |
(n_layers) | layers in a stack | int |
| segments in a schedule | int | |
| layers in segment (), | int | |
starts | first layer of each segment: | list[int] |
(mem_budget_layers) | the memory budget, in saved layer inputs | int |
| the peak number of saved layer inputs held at once | int |
The Hessian-vector product is a directional derivative of the gradient. Taylor-expand the gradient around along :
The first-order term is , so . That limit is a derivative, and you already know how to take derivatives numerically.
Central differences cancel the even terms. Subtract the expansion at from the one at : and the term cancel, and
The one-sided quotient keeps the term: halving halves its error, while halving it in the central quotient quarters the error. The step cannot be made arbitrarily small either: each gradient carries rounding error of about ( in float64), and dividing by magnifies it to about . Balancing against puts the best step near times the scale of the problem; the default gives about relative accuracy on smooth losses and is far from the rounding floor.
Quadratics are exact. For , is affine, , and the central difference returns for every (up to rounding). A non-symmetric shows that the Hessian is the symmetric part.
Cost. Two gradient calls, each about one forward and one backward pass, for one . Forming would take of them and numbers of memory: for a 10M-parameter model, entries. Everything that needs curvature uses products instead: Newton-CG solves with , and the power iteration finds the top eigenvalue (M10.5). For small , hessian_fd builds column by column from and returns : the columns carry errors that differ between and , and code that assumes symmetry (eigenvalues, Cholesky) must get it exactly. Pearlmutter’s -operator (forward mode over reverse mode) computes exactly at the same cost, which is what frameworks do; it needs an autograd that differentiates its own backward pass, which yours does not.
Reverse mode stores activations. A stack of layers , , computes the VJP of layer from its saved input (a matmul’s needs ). Count memory in units of one saved layer input. A plain forward pass keeps all , and backward frees them from the top down: the peak is .
Segments and recomputation. Split the stack into consecutive segments of layers. Every segment but the last runs forward keeping only its input (one unit) and discards its internal activations; the last segment runs normally and keeps its . During backward, segment is rerun from its saved input to rebuild its activations, then backpropagated and freed. This is torch.utils.checkpoint.checkpoint_sequential. Memory while segment is backpropagated: the inputs of segments (still needed later) plus segment ‘s activations (the first of which is its saved input). The end of the forward pass is the case . So
Both extremes cost : one segment () is no checkpointing, and one layer per segment holds inputs while recomputing layers for nothing.
Uniform segments give . With segments of layers (), . Setting the derivative to zero gives and : the classic sublinear-memory result of Chen et al.
Shrinking segments do better. The term grows with , so later segments should be smaller. If every , then , and segments hold at most
layers. That is largest at (the segment sizes ), where it is the triangular number . So the least peak of any schedule is the least with , about : for , against 7 for uniform segments of 4. Compute it with integers (math.isqrt): a float square root is wrong by one for large .
The budgeted schedule. Given a budget , minimize recomputation, that is, maximize the last segment. With segments the last holds at most layers (its , and every other segment needs at least one), and the other segments must cover the rest within . That bound shrinks as grows, so the fewest feasible segments is optimal. Fill the first segments greedily, largest first, , leaving at least one layer for each later segment; the order fixes one answer among ties, so your trainer and the reference checkpoint the same layers. A brute force over all schedules agrees for every (test_schedule_is_optimal).
3. Worked example by hand
Section titled “3. Worked example by hand”A quartic. in one dimension: , , , and . At , , :
| quantity | value |
|---|---|
| , | , |
| central: | |
| exact | |
| error, predicted | (exact here: is cubic, so the series stops) |
| one-sided: | , error , thirty times larger |
| forgetting the 2: |
A six-layer stack. :
schedule (starts) | sizes | peak | recomputed | |
|---|---|---|---|---|
[0] | 6 | 6 | 6 | 0 |
[0, 1, 2, 3, 4, 5] | 1, 1, 1, 1, 1, 1 | 1, 2, 3, 4, 5, 6 | 6 | 5 |
[0, 2, 4] (uniform, ) | 2, 2, 2 | 2, 3, 4 | 4 | 4 |
[0, 3, 5] | 3, 2, 1 | 3, 3, 3 | 3 | 5 |
[0, 3] | 3, 3 | 3, 4 | 4 | 3 |
The least peak is 3, because , and [0, 3, 5] reaches it. With a budget of 4: one segment needs 6 units, too many; two segments can keep a last segment of layers if the first covers the other 3 within , which it does. So checkpoint_schedule(6, 4) == [0, 3]: the same peak as uniform segments and one recomputed layer fewer. These numbers are the first two tests, test_hand_example_quartic and test_hand_example_schedule.
4. The interface
Section titled “4. The interface”def hvp_fd(grad_fn, x, v, eps: float = 1e-4) -> NDArray # (g(x + eps v) - g(x - eps v)) / (2 eps)def hessian_fd(grad_fn, x, eps: float = 1e-4) -> NDArray # 1-D x: columns H e_j, then (H + H^T) / 2def checkpoint_cost(n_layers: int, starts) -> tuple[int, int] # (peak, recomputed)def min_checkpoint_memory(n_layers: int) -> int # least P with P (P + 1) / 2 >= ndef checkpoint_schedule(n_layers: int, mem_budget_layers: int) -> list[int] # segment startshvp_fd calls grad_fn exactly twice, each time with a new float64 array, and copies each result, so a gradient that returns its own argument is safe. The step is , not normalized: scale yourself if its norm is far from 1. Shape mismatches, a non-positive or non-finite step, and an infeasible budget raise ValueError.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example_quartic | unit | the section 3 quartic numbers, error | you and the test agree on the formula |
test_hand_example_schedule | unit | the section 3 schedule table | L11.1 checkpoints exactly these layers |
test_quadratic_hvp_is_exact | property | at three step sizes, not unit | the step is |
test_error_shrinks_as_eps_squared | property | halving divides the error by 4 | a central difference, not one-sided |
test_matches_torch_hvp | golden | torch’s exact double-backward of a cross-entropy loss whose gradient is M08.3’s rule | the curvature M10.5 estimates |
test_hessian_fd_matches_torch | golden | a dense Hessian, exactly symmetric | small-model Newton steps |
test_hessian_is_exactly_symmetric | unit | Rosenbrock at : bit for bit, close to the analytic Hessian | eigenvalue code assumes symmetry |
test_grad_fn_may_return_its_argument | unit | a gradient that returns its input; two calls, x untouched | cheap gradients alias |
test_hvp_rejects_bad_inputs | boundary | shape and step errors raise; the default step is accurate | caller bugs surface |
test_checkpoint_cost_extremes | unit | no checkpointing and all checkpoints both peak at ; invalid schedules raise | the cost model itself |
test_min_memory_is_triangular | unit | up to | integer arithmetic at scale |
test_schedule_is_optimal | property | within budget and as few recomputed layers as brute force, | the schedule is the best one |
test_schedule_ties_go_to_the_largest_first_segment | unit | the tie rule on four cases | reference and learner agree |
test_schedule_rejects_impossible_budget | boundary | a budget below the least peak raises | no silent overrun |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. a one-sided difference | error : 0.5517 instead of 0.5017 in section 3 | test_hand_example_quartic (mutant s01) |
| 2. dividing by instead of | every twice too large | test_hand_example_quartic (mutant s02) |
| 3. reusing one buffer for | a gradient that aliases its input changes under you; | test_grad_fn_may_return_its_argument (mutant s03) |
| 4. not symmetrizing the finite-difference Hessian | by | test_hessian_is_exactly_symmetric (mutant s04) |
| 5. peak as (segments 1) + (largest segment) | overcounts when the largest segment comes first; the optimizer picks worse schedules | test_hand_example_schedule (mutant s05) |
| 6. the uniform rule whatever the budget | over budget when it is tight, needless recomputation when it is loose | test_hand_example_schedule (mutant s09) |
7. a float ceil(sqrt(2n)) for the least peak | one too many at , wrong at large | test_min_memory_is_triangular (mutant s07) |
| 8. counting the last segment as recomputed | recomputation off by | test_checkpoint_cost_extremes (mutant s06) |
| 9. a budget check off by one | peaks one unit over the budget | test_schedule_rejects_impossible_budget (mutant s10) |
10. normalizing inside hvp_fd | returns | test_quadratic_hvp_is_exact (mutant s12) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | M08.3 | its cross-entropy VJP, , is the gradient whose the golden tests compare with torch |
| Back | S-M08 | the forward versus reverse cost counts behind “two gradients per ” |
| Back | M04.2 | the Hessian and the second-derivative test |
| Forward | L11.1 | checkpoint_sequential runs your segments with checkpoint_schedule and recomputes them in backward |
| Forward | M10.5 | (optional) power iteration on hvp_fd estimates to compare the learning rate with |
If you skip this module, ol check L11.1 stops with L11.1 needs M08.4: build it, or rerun with --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
hvp_fd | torch.autograd.functional.hvp, torch.func.jvp(grad(f)) | exact forward-over-reverse products (Pearlmutter’s -operator) at the cost of two backward passes | torch/autograd/functional.py |
hessian_fd | PyHessian | Hessian spectra of real networks by power iteration and stochastic Lanczos, from only | pyhessian/hessian.py |
checkpoint_schedule | torch.utils.checkpoint.checkpoint_sequential | uniform segments, RNG state replay, non-reentrant saved-tensor hooks | torch/utils/checkpoint.py |
| the cost model | Checkmate (Jain et al., 2020) | optimal rematerialization for arbitrary graphs as an integer program | checkmate repository, checkmate/core/solvers/ |