Skip to content

Hessian-vector products and the recompute-versus-memory schedule

ModuleM08.4 · build · Python · Pass 9 · 3 h
You buildpython/tinyllm/autograd/hvp.py: hvp_fd, hessian_fd, checkpoint_cost, min_checkpoint_memory, checkpoint_schedule
Contractcourse/contracts/py/tinyllm/autograd/hvp.pyi
Testscourse/tests/M08.4/ (what they check: section 4)
Needsnothing 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 byL11.1 activation checkpointing plans its segments with checkpoint_schedule · later: M10.5 (optional) estimates the top Hessian eigenvalue with hvp_fd
MilestoneMS-P9 (Pass 9 gate: the math of the pass checks green, then the capstone trains with checkpointing)
Optional depthPearlmutter, “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)
  • The Hessian-vector product HvHv is the derivative of the gradient along vv, so two gradient calls give it to O(ϵ2)O(\epsilon^2) without forming the n×nn \times n 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 ϵ\epsilon, 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 max⁡i(i+si)\max_i (i + s_i) and the extra work is every layer outside the last segment (test_checkpoint_cost_extremes).
  • Uniform segments of n\sqrt n layers peak near 2n2\sqrt n; shrinking segments (sizes P,P−1,…,1P, P-1, \dots, 1) reach the optimum P≈2nP \approx \sqrt{2n}, and under a budget the best schedule is found greedily (test_min_memory_is_triangular, test_schedule_is_optimal).
Terminal window
ol start M08.4 # stubs hvp.py into your repo, contract alongside
ol tests M08.4 # read the test catalog first: rung R0, you write no tests here
ol check M08.4 # exit code is the verdict
ol diff M08.4 # after passing: your code against the reference

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 2/λmax⁡2/\lambda_{\max}?”, 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.

SymbolMeaningType / shape
f:Rn→Rf: \mathbb{R}^n \to \mathbb{R}a scalar loss of nn parametersfunction
x∈Rnx \in \mathbb{R}^nthe point (parameters)float64[...]
g(x)=∇f(x)g(x) = \nabla f(x)the gradient (grad_fn)same shape as xx
H(x)=∇2f(x)H(x) = \nabla^2 f(x)the Hessian: Hij=∂2f/∂xi∂xjH_{ij} = \partial^2 f / \partial x_i \partial x_jn×nn \times n, symmetric
vva directionsame shape as xx
ϵ\epsilonthe finite-difference step along vv (eps)float
Dkg(x)[v,…,v]D^k g(x)[v, \dots, v]the kk-th directional derivative of gg along vv (a derivative of order k+1k + 1 of ff)same shape as xx
eje_jthe jj-th unit vectorfloat64[n]
nn (n_layers)layers in a stackint
kksegments in a scheduleint
sis_ilayers in segment ii (i=0,…,k−1i = 0, \dots, k-1), ∑isi=n\sum_i s_i = nint
startsfirst layer of each segment: 0,s0,s0+s1,…0, s_0, s_0 + s_1, \dotslist[int]
BB (mem_budget_layers)the memory budget, in saved layer inputsint
PPthe peak number of saved layer inputs held at onceint

The Hessian-vector product is a directional derivative of the gradient. Taylor-expand the gradient around xx along vv:

g(x+ϵv)=g(x)+ϵ Hv+ϵ22D2g(x)[v,v]+ϵ36D3g(x)[v,v,v]+O(ϵ4).g(x + \epsilon v) = g(x) + \epsilon\, H v + \tfrac{\epsilon^2}{2} D^2 g(x)[v, v] + \tfrac{\epsilon^3}{6} D^3 g(x)[v, v, v] + O(\epsilon^4).

The first-order term is ϵHv\epsilon H v, so Hv=lim⁡ϵ→0(g(x+ϵv)−g(x))/ϵHv = \lim_{\epsilon \to 0} (g(x + \epsilon v) - g(x))/\epsilon. That limit is a derivative, and you already know how to take derivatives numerically.

Central differences cancel the even terms. Subtract the expansion at −ϵ-\epsilon from the one at +ϵ+\epsilon: g(x)g(x) and the ϵ2\epsilon^2 term cancel, and

g(x+ϵv)−g(x−ϵv)2ϵ=Hv+ϵ26D3g(x)[v,v,v]+O(ϵ4).\frac{g(x + \epsilon v) - g(x - \epsilon v)}{2\epsilon} = Hv + \frac{\epsilon^2}{6} D^3 g(x)[v, v, v] + O(\epsilon^4).

The one-sided quotient keeps the ϵ2D2g(x)[v,v]\tfrac{\epsilon}{2} D^2 g(x)[v,v] term: halving ϵ\epsilon 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 u ∣g∣u\,|g| (u=2−53≈1.1×10−16u = 2^{-53} \approx 1.1 \times 10^{-16} in float64), and dividing by 2ϵ2\epsilon magnifies it to about u ∣g∣/ϵu\,|g|/\epsilon. Balancing ϵ2\epsilon^2 against u/ϵu/\epsilon puts the best step near u1/3≈5×10−6u^{1/3} \approx 5 \times 10^{-6} times the scale of the problem; the default 10−410^{-4} gives about 10−810^{-8} relative accuracy on smooth losses and is far from the rounding floor.

Quadratics are exact. For f(x)=12x⊤Ax+b⊤xf(x) = \tfrac12 x^\top A x + b^\top x, g(x)=12(A+A⊤)x+bg(x) = \tfrac12(A + A^\top) x + b is affine, D3g=0D^3 g = 0, and the central difference returns 12(A+A⊤)v\tfrac12 (A + A^\top) v for every ϵ\epsilon (up to rounding). A non-symmetric AA shows that the Hessian is the symmetric part.

Cost. Two gradient calls, each about one forward and one backward pass, for one HvHv. Forming HH would take nn of them and n2n^2 numbers of memory: for a 10M-parameter model, 101410^{14} entries. Everything that needs curvature uses products instead: Newton-CG solves with HvHv, and the power iteration v←Hv/∥Hv∥v \leftarrow Hv / \lVert Hv \rVert finds the top eigenvalue λmax⁡\lambda_{\max} (M10.5). For small nn, hessian_fd builds HH column by column from HejH e_j and returns (H+H⊤)/2(H + H^\top)/2: the columns carry O(ϵ2)O(\epsilon^2) errors that differ between HijH_{ij} and HjiH_{ji}, and code that assumes symmetry (eigenvalues, Cholesky) must get it exactly. Pearlmutter’s R\mathcal{R}-operator (forward mode over reverse mode) computes HvHv 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 yl+1=fl(yl)y_{l+1} = f_l(y_l), l=0,…,n−1l = 0, \dots, n-1, computes the VJP of layer ll from its saved input yly_l (a matmul’s Wˉ=yˉ⊤x\bar W = \bar y^\top x needs xx). Count memory in units of one saved layer input. A plain forward pass keeps all nn, and backward frees them from the top down: the peak is nn.

Segments and recomputation. Split the stack into kk consecutive segments of s0,…,sk−1s_0, \dots, s_{k-1} 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 sk−1s_{k-1}. During backward, segment ii is rerun from its saved input to rebuild its sis_i activations, then backpropagated and freed. This is torch.utils.checkpoint.checkpoint_sequential. Memory while segment ii is backpropagated: the inputs of segments 0,…,i−10, \dots, i-1 (still needed later) plus segment ii‘s sis_i activations (the first of which is its saved input). The end of the forward pass is the case i=k−1i = k - 1. So

P=max⁡0≤i<k(i+si),recomputed layers=n−sk−1.P = \max_{0 \le i < k} (i + s_i), \qquad \text{recomputed layers} = n - s_{k-1}.

Both extremes cost nn: one segment (k=1k = 1) is no checkpointing, and one layer per segment holds nn inputs while recomputing n−1n - 1 layers for nothing.

Uniform segments give n\sqrt n. With kk segments of ss layers (ks=nks = n), P=(k−1)+s=n/s+s−1P = (k - 1) + s = n/s + s - 1. Setting the derivative −n/s2+1-n/s^2 + 1 to zero gives s=ns = \sqrt n and P≈2n−1P \approx 2\sqrt n - 1: the classic sublinear-memory result of Chen et al.

Shrinking segments do better. The term i+sii + s_i grows with ii, so later segments should be smaller. If every i+si≤Pi + s_i \le P, then si≤P−is_i \le P - i, and kk segments hold at most

C(k,P)=∑i=0k−1(P−i)=kP−k(k−1)2C(k, P) = \sum_{i=0}^{k-1} (P - i) = kP - \frac{k(k-1)}{2}

layers. That is largest at k=Pk = P (the segment sizes P,P−1,…,1P, P-1, \dots, 1), where it is the triangular number P(P+1)/2P(P+1)/2. So the least peak of any schedule is the least PP with P(P+1)/2≥nP(P+1)/2 \ge n, about 2n\sqrt{2n}: for n=16n = 16, P=6P = 6 against 7 for uniform segments of 4. Compute it with integers (math.isqrt): a float square root is wrong by one for large nn.

The budgeted schedule. Given a budget B≥Pmin⁡B \ge P_{\min}, minimize recomputation, that is, maximize the last segment. With kk segments the last holds at most min⁡(B−k+1, n−k+1)\min(B - k + 1,\, n - k + 1) layers (its i+si≤Bi + s_i \le B, and every other segment needs at least one), and the other k−1k-1 segments must cover the rest within C(k−1,B)C(k-1, B). That bound shrinks as kk grows, so the fewest feasible segments is optimal. Fill the first k−1k - 1 segments greedily, largest first, si=min⁡(B−i, rest−(k−2−i))s_i = \min(B - i,\ \text{rest} - (k - 2 - i)), 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 2n−12^{n-1} schedules agrees for every n≤10n \le 10 (test_schedule_is_optimal).

A quartic. f(x)=x4/24f(x) = x^4/24 in one dimension: g(x)=x3/6g(x) = x^3/6, H(x)=g′(x)=x2/2H(x) = g'(x) = x^2/2, g′′(x)=xg''(x) = x, and g′′′(x)=1g'''(x) = 1. At x=1x = 1, v=1v = 1, ϵ=0.1\epsilon = 0.1:

quantityvalue
g(1.1)=1.331/6g(1.1) = 1.331/6, g(0.9)=0.729/6g(0.9) = 0.729/60.22183330.2218333, 0.12150.1215
central: (1.331−0.729)/(6⋅0.2)(1.331 - 0.729)/(6 \cdot 0.2)0.602/1.2=0.50166670.602/1.2 = 0.5016667
exact HvH v0.50.5
error, predicted ϵ2v3g′′′/6\epsilon^2 v^3 g'''/60.01/6=0.00166670.01/6 = 0.0016667 (exact here: gg is cubic, so the series stops)
one-sided: (1.331−1)/(6⋅0.1)(1.331 - 1)/(6 \cdot 0.1)0.55166670.5516667, error 0.0517≈ϵ2g′′(1)0.0517 \approx \tfrac{\epsilon}{2} g''(1), thirty times larger
forgetting the 2: 0.602/0.60.602/0.61.00331.0033

A six-layer stack. n=6n = 6:

schedule (starts)sizes sis_ii+sii + s_ipeak PPrecomputed
[0]6660
[0, 1, 2, 3, 4, 5]1, 1, 1, 1, 1, 11, 2, 3, 4, 5, 665
[0, 2, 4] (uniform, s≈6s \approx \sqrt 6)2, 2, 22, 3, 444
[0, 3, 5]3, 2, 13, 3, 335
[0, 3]3, 33, 443

The least peak is 3, because 3⋅4/2=6≥6>2⋅3/23 \cdot 4/2 = 6 \ge 6 > 2 \cdot 3/2, 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 min⁡(4−1,6−1)=3\min(4 - 1, 6 - 1) = 3 layers if the first covers the other 3 within C(1,4)=4C(1, 4) = 4, 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.

python/tinyllm/autograd/hvp.py
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) / 2
def 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 >= n
def checkpoint_schedule(n_layers: int, mem_budget_layers: int) -> list[int] # segment starts

hvp_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 ϵv\epsilon v, not normalized: scale vv yourself if its norm is far from 1. Shape mismatches, a non-positive or non-finite step, and an infeasible budget raise ValueError.

TestKINDChecksWhy it matters downstream
test_hand_example_quarticunitthe section 3 quartic numbers, error ϵ2/6\epsilon^2/6you and the test agree on the formula
test_hand_example_scheduleunitthe section 3 schedule tableL11.1 checkpoints exactly these layers
test_quadratic_hvp_is_exactproperty12(A+A⊤)v\tfrac12 (A + A^\top) v at three step sizes, vv not unitthe step is ϵv\epsilon v
test_error_shrinks_as_eps_squaredpropertyhalving ϵ\epsilon divides the error by 4a central difference, not one-sided
test_matches_torch_hvpgoldentorch’s exact double-backward HvHv of a cross-entropy loss whose gradient is M08.3’s rulethe curvature M10.5 estimates
test_hessian_fd_matches_torchgoldena dense 6×66 \times 6 Hessian, exactly symmetricsmall-model Newton steps
test_hessian_is_exactly_symmetricunitRosenbrock at (−1.2,1)(-1.2, 1): H=H⊤H = H^\top bit for bit, close to the analytic Hessianeigenvalue code assumes symmetry
test_grad_fn_may_return_its_argumentunita gradient that returns its input; two calls, x untouchedcheap gradients alias
test_hvp_rejects_bad_inputsboundaryshape and step errors raise; the default step is accuratecaller bugs surface
test_checkpoint_cost_extremesunitno checkpointing and all checkpoints both peak at nn; invalid schedules raisethe cost model itself
test_min_memory_is_triangularunitP(P+1)/2≥n>(P−1)P/2P(P+1)/2 \ge n > (P-1)P/2 up to n=261−1n = 2^{61} - 1integer arithmetic at scale
test_schedule_is_optimalpropertywithin budget and as few recomputed layers as brute force, n≤10n \le 10the schedule is the best one
test_schedule_ties_go_to_the_largest_first_segmentunitthe tie rule on four casesreference and learner agree
test_schedule_rejects_impossible_budgetboundarya budget below the least peak raisesno silent overrun
PitfallSymptomCaught by
1. a one-sided differenceerror O(ϵ)O(\epsilon): 0.5517 instead of 0.5017 in section 3test_hand_example_quartic (mutant s01)
2. dividing by ϵ\epsilon instead of 2ϵ2\epsilonevery HvHv twice too largetest_hand_example_quartic (mutant s02)
3. reusing one buffer for x±ϵvx \pm \epsilon va gradient that aliases its input changes under you; Hv=0Hv = 0test_grad_fn_may_return_its_argument (mutant s03)
4. not symmetrizing the finite-difference HessianHij≠HjiH_{ij} \ne H_{ji} by O(ϵ2)O(\epsilon^2)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 schedulestest_hand_example_schedule (mutant s05)
6. the uniform n\sqrt n rule whatever the budgetover budget when it is tight, needless recomputation when it is loosetest_hand_example_schedule (mutant s09)
7. a float ceil(sqrt(2n)) for the least peakone too many at n=6n = 6, wrong at large nntest_min_memory_is_triangular (mutant s07)
8. counting the last segment as recomputedrecomputation off by sk−1s_{k-1}test_checkpoint_cost_extremes (mutant s06)
9. a budget check off by onepeaks one unit over the budgettest_schedule_rejects_impossible_budget (mutant s10)
10. normalizing vv inside hvp_fdreturns Hv/∥v∥Hv/\lVert v \rVerttest_quadratic_hvp_is_exact (mutant s12)
DirectionModuleHow it uses this
BackM08.3its cross-entropy VJP, X⊤(softmax−onehot)/nX^\top(\mathrm{softmax} - \mathrm{onehot})/n, is the gradient whose HvHv the golden tests compare with torch
BackS-M08the forward versus reverse cost counts behind “two gradients per HvHv”
BackM04.2the Hessian and the second-derivative test
ForwardL11.1checkpoint_sequential runs your segments with checkpoint_schedule and recomputes them in backward
ForwardM10.5(optional) power iteration on hvp_fd estimates λmax⁡\lambda_{\max} to compare the learning rate with 2/λmax⁡2/\lambda_{\max}

If you skip this module, ol check L11.1 stops with L11.1 needs M08.4: build it, or rerun with --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
hvp_fdtorch.autograd.functional.hvp, torch.func.jvp(grad(f))exact forward-over-reverse products (Pearlmutter’s R\mathcal{R}-operator) at the cost of two backward passestorch/autograd/functional.py
hessian_fdPyHessianHessian spectra of real networks by power iteration and stochastic Lanczos, from HvHv onlypyhessian/hessian.py
checkpoint_scheduletorch.utils.checkpoint.checkpoint_sequentialuniform segments, RNG state replay, non-reentrant saved-tensor hookstorch/utils/checkpoint.py
the cost modelCheckmate (Jain et al., 2020)optimal rematerialization for arbitrary graphs as an integer programcheckmate repository, checkmate/core/solvers/