Jacobians, vjp and jvp with torch.func
When one number's gradient is not enough, torch.func computes whole derivative tables — or single rows and columns of them — as clean function transforms.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A Jacobian is the full table of how every output reacts to every input — and torch.func can compute the table, or any single row or column of it.
Think of a mixing desk in a music studio: rows of sliders, several speakers. Every slider affects every speaker a little differently. The full sensitivity chart — every slider against every speaker — is a big table. Often you do not need the whole chart. "How do all sliders affect the left speaker?" is one row. "What does moving this slider do everywhere?" is one column.
That table is the Jacobian. The row and column questions have short names too: vjp and jvp.
Why it exists
Ordinary training never needs the table. The loss is one number, so one row covers it — that is why backward() insists on a scalar.
But past ordinary training, table questions appear. How sensitive is each of a robot's joint positions to each motor? How does a physics model's every output respond to every knob? Researchers probing how networks learn need rows, columns, sometimes everything. torch.func answers these as clean one-liners, in a style where you transform a function instead of walking one recording.
How it works
inputs (n sliders)
|
[ function f ]
|
outputs (m speakers)
Jacobian: the full m x n table
vjp: one weighted ROW combination (backward-style question)
jvp: one COLUMN (forward-style question)Rows come cheap when outputs are few. Columns come cheap when inputs are few. The whole table costs many passes either way — so the tools let you buy only what you need.
A real example you have seen
Weather bulletins saying "if the sea stays warm, expect heavier rain inland" are reading one column of an enormous sensitivity table — one input nudged, all outputs watched. Forecasting centres compute such sensitivities from simulation models routinely. Same mathematics, bigger desk.
Remember this
- The Jacobian is the every-output-versus-every-input sensitivity table.
- vjp buys one row-combination; jvp buys one column.
- Ordinary training needs one row; research and simulation need the rest.
What to learn next
- Calling backward on a non-scalar output — the vᵀJ story this lesson generalises.
- Checking your gradients with gradcheck — numerical Jacobians as the referee.
- Derivatives and gradients — the calculus underneath the tables.
Developer — Code and libraries.
Setup
pip install torchOutputs verified with torch 2.5.1, CPU. torch.func ships inside torch 2.x — nothing extra to install.
The table, a row, and a column
import torch
from torch.func import jacrev, jacfwd, vjp, jvp
def f(x):
return torch.stack([x[0] * x[1], # depends on both inputs
x[0] ** 2, # depends on x[0] only
x[1] + 3]) # depends on x[1] only
x = torch.tensor([2., 5.])
J = jacrev(f)(x) # the full 3x2 table of derivatives
print(J)
print(torch.allclose(J, jacfwd(f)(x)))
out, pull = vjp(f, x) # reverse mode: rows of J
print(pull(torch.tensor([1., 0., 0.]))[0]) # row 0, no full J built
_, push = jvp(f, (x,), (torch.tensor([1., 0.]),)) # forward mode: columns
print(push) # column 0, no full J builttensor([[5., 2.],
[4., 0.],
[0., 1.]])
True
tensor([5., 2.])
tensor([5., 4., 0.])The walkthrough
Check the table by hand once — it cements everything. Output 0 is x₀·x₁: its sensitivities are (x₁, x₀) = (5, 2), the first row. Output 1 is x₀²: sensitivities (2x₀, 0) = (4, 0). Output 2 touches only x₁: (0, 1). Three rows, two columns, every entry explainable.
jacrev versus jacfwd is a cost choice, not a maths choice — they agree (the allclose proves it) and differ in how many passes they spend. jacrev works row-at-a-time: cheap for few outputs and many inputs, the deep-learning shape. jacfwd works column-at-a-time: cheap for few inputs and many outputs. Tall table → jacfwd; wide table → jacrev.
vjp returned the function pull. Calling pull(v) answers "v-weighted combination of rows" — seeding with [1., 0., 0.] extracts row 0, matching the table's first row. This is exactly what backward(v) computes, packaged functionally: no .grad fields touched, no graph consumed, call pull as many times as you like.
jvp pushed the column out directly. Direction [1., 0.] means "nudge input 0 alone": out came column 0, [5., 4., 0.]. Forward-mode has no equivalent in classic autograd — this is new capability, not repackaging.
The transform style composes. grad, the functional cousin of backward, stacks with vmap to give the per-sample gradients that classic autograd sums away:
import torch
from torch.func import grad, vmap
def loss_one(x):
return (x * torch.tensor([2., 1.])).sum()
batch = torch.ones(4, 2)
per_sample = vmap(grad(loss_one))(batch) # one gradient PER row
print(per_sample.shape)torch.Size([4, 2])
Per-sample gradients power differential privacy and influence analysis, and pre-torch.func they required painful loops or hooks.
Common mistakes
Computing the full Jacobian to use one row. If your next line multiplies J by a vector, you wanted vjp or jvp — a single pass instead of many. The full table is for when you genuinely inspect it.
Feeding jacrev a batched function and misreading the result. For f over a batch, the naive Jacobian includes cross-sample zeros — a huge sparse table. Per-sample tables are vmap(jacrev(f)), as with grad above.
Expecting x.grad to appear. torch.func transforms are pure: results come back as return values, nothing is written anywhere. Mixing the two styles — calling backward inside a transformed function — raises errors; within a transform, stay functional.
Side effects inside transformed functions. Printing, appending to lists, or mutating globals inside f behaves surprisingly under vmap (it runs on batched abstractions, not per-element loops). Keep transformed functions pure — inputs to outputs, nothing else.
Forgetting gradcheck's day job is this. gradcheck builds these same Jacobians numerically. When a transform's output confuses you, a tiny numerical check settles it in five lines.
Try it yourself
For f above, get row 2 via pull and column 1 via jvp, predicting both before running. Then time jacrev versus jacfwd on g(x) = torch.sin(x).sum() with x of size 1000 — a 1×1000 table — and explain the winner using rows-versus-columns.
What to learn next
- Calling backward on a non-scalar output — the vᵀJ story this lesson generalises.
- Checking your gradients with gradcheck — numerical Jacobians as the referee.
- Derivatives and gradients — the calculus underneath the tables.
Researcher — Mathematics and papers.
Definitions and duality
For f: ℝⁿ → ℝᵐ with Jacobian J ∈ ℝ^{m×n} at x:
- vjp: v ∈ ℝᵐ ↦ vᵀJ ∈ ℝⁿ — reverse mode; one backward-style pass; cost O(cost(f)).
- jvp: u ∈ ℝⁿ ↦ Ju ∈ ℝᵐ — forward mode, computed by dual-number/tangent propagation alongside the forward pass; also O(cost(f)), with no tape stored.
- Full J: m vjp passes (jacrev: rows via basis seeds, batched by vmap) or n jvp passes (jacfwd: columns). Choose by min(m, n) — the classical reverse/forward crossover (Griewank and Walther, 2008; survey: Baydin et al., 2018, JMLR).
Memory profiles differ asymmetrically: reverse mode stores activations O(depth) for the transposed sweep; forward mode carries tangents O(1) extra per value. This makes jvp attractive even at m ≈ n when memory, not FLOPs, binds.
Compositions worth knowing
- Hessian: for scalar F,
jacfwd(jacrev(F))— forward-over-reverse is the standard efficient shape: Hv products come from one jvp of the gradient function, cost O(cost(F)) per column (Pearlmutter, 1994, Fast exact multiplication by the Hessian).torch.func.hessianpackages it. - HVP without materialising H:
jvp(grad(F), (x,), (v,))— the workhorse of curvature analysis, influence functions (Koh and Liang, 2017), K-FAC-adjacent methods and second-order optimiser research. - Per-sample gradients:
vmap(grad(f))vectorises the m = 1 case over a batch axis — required exactly by DP-SGD's per-sample clipping (Abadi et al., 2016) and gradient-noise diagnostics. - Forward-mode-only training research: forward-gradient methods estimate ∇F from jvps with random tangents, E[u(Ju)ᵀ] = J for suitable u (Baydin et al., 2022, Gradients without Backpropagation) — practical relevance still contested, machinery identical.
The transform machinery
torch.func descends from functorch, itself modelled on JAX's composable transforms (Bradbury et al., 2018, jax GitHub) — grad/vmap/jvp/vjp as higher-order functions over pure functions, in contrast with the tape-and-mutation style of backward(). Stabilised inside torch during 2.0–2.1 and steady through 2.5; the migration note (functorch → torch.func) is in the official docs. Under the hood, vmap interprets operations on batched tensors with a lifted batch dimension — the reason side-effectful functions misbehave — and custom Functions require the setup_context signature to be transform-compatible (see custom Functions).
For model code with parameters rather than explicit arguments, torch.func.functional_call(model, params, x) re-expresses a stateful nn.Module as a pure function of (params, input) — the bridge that lets jacrev differentiate with respect to parameter pytrees.
References
- Griewank and Walther (2008), Evaluating Derivatives, SIAM — modes, duality, complexity.
- Baydin et al. (2018), Automatic Differentiation in Machine Learning: a Survey, JMLR 18(153).
- Pearlmutter (1994), Neural Computation 6(1) — Hessian-vector products.
- Abadi et al. (2016), CCS — per-sample gradients in anger.
- PyTorch docs, torch.func — API, composition rules, functorch migration.
What to learn next
- Calling backward on a non-scalar output — the vᵀJ story this lesson generalises.
- Checking your gradients with gradcheck — numerical Jacobians as the referee.
- Derivatives and gradients — the calculus underneath the tables.