Higher-Order Gradients & Jacobians¶
Overview¶
First-order training (loss.backward()) is the 99% case. But meta-learning, adversarial robustness, scientific ML (PDE solvers), and deep equilibrium models all need Jacobians, Hessians, and gradient-of-gradient. PyTorch gives you two tools: autograd with create_graph=True, and the torch.func (functorch) functional API.
torch.autograd.grad(loss, params, create_graph=True)— first forward, second derivative via chains.torch.func.jacrev/jacfwd/jvp/vjp/hessian— explicit Jacobians for any callable.- Hessian-vector products (HVP) are cheap (
jvp(vjp(...))) — never build the full Hessian. - Double backward requires the graph to be kept:
retain_graph=Trueor recompute.
💡 Rule: for anything requiring a full Jacobian, use
torch.func/vmap; for HVP/GN updates, composejvpandvjpsymbolically.
autograd.grad with create_graph¶
import torch
x = torch.tensor([2.0, 3.0], requires_grad=True)
loss = (x ** 3).sum()
g1 = torch.autograd.grad(loss, x, create_graph=True)[0] # 3x^2
g2 = torch.autograd.grad(g1.sum(), x)[0] # 6x
print("first:", g1.tolist(), " second:", g2.tolist())
retain_graph vs create_graph¶
retain_graph=True→ reuse the same graph for another backward (memory heavy).create_graph=True→ build another graph that will itself be backwarded later.- Prefer
autograd.gradfor selective params;backward(create_graph=True)for everything.
Vector-Jacobian Products (VJP) — the natural primitive¶
vjp(f, inputs, v) computes v^T J — exactly what autograd does, exposed directly.
from torch.func import vjp, jvp
def f(x):
return x.sin().sum()
x = torch.tensor([1.0, 2.0], requires_grad=True)
v = torch.tensor([1.0, 1.0])
out, vjp_fn = vjp(f, x)
grad = vjp_fn(v)[0]
print("vjp (grad):", grad) # sum of cos
Full Jacobians — jacrev / jacfwd¶
from torch.func import jacrev, jacfwd, vmap
def vec_fun(x):
return torch.stack([x[0] ** 2, x[0] * x[1]])
x = torch.tensor([2.0, 3.0])
J = jacrev(vec_fun)(x) # 2x2
print(J)
# [[4., 0.], [3., 2.]]
| API | Cost | Use when |
|---|---|---|
jacrev |
output-dim backwards | output small, input large |
jacfwd |
input-dim forwards | input small, output large |
jacrev(jacrev(...)) |
Hessian | dims tiny |
Hessian-Vector Products Without the Hessian¶
from torch.func import jvp, vjp
def hessian_vector_product(f, x, v):
_, vjp_fn = vjp(f, x)
_, jvp_fn = jvp(lambda z: vjp_fn(z)[0], (x,), (v,))
return jvp_fn(x)[0] # careful: use correct composition
def g(z):
return (z ** 3).sum()
x = torch.tensor([2.0])
Hv = hessian_vector_product(g, x, torch.tensor([1.0]))
print("Hv (should be 6x*1 = 12):", Hv)
💡 Cost of HVP ≈ 2 forwards — this is why second-order optimizers (Newton/L-BFGS/K-FAC-ish) are feasible at scale.
vmap — Batch Without a For-Loop¶
from torch.func import vmap
def single(x):
return (x @ torch.ones(3)).sum()
xs = torch.randn(10, 3)
outs = vmap(single)(xs) # implicit batch dim
print(outs.shape)
Use case: per-sample gradients, per-sample losses, Jacobians-of-batches.
params = [torch.randn(3, 3, requires_grad=True)]
def loss_fn(p, x):
return (x @ p).square().mean()
# per-sample grads in one call:
batch_grads = vmap(lambda x_: torch.autograd.grad(loss_fn(params[0], x_), params[0])[0])(xs)
print("per-sample grads shape:", batch_grads.shape) # (10, 3, 3)
Practical Guards¶
import torch
def safe_second_order(f, x, v=None):
"""Stable Hessian-vector product with NaNs guarded."""
g1 = torch.autograd.grad(f(x).sum(), x, create_graph=True)[0]
if v is None:
v = torch.ones_like(x)
g2 = torch.autograd.grad((g1 * v).sum(), x)[0]
assert torch.isfinite(g2).all()
return g2
Decision Table¶
| Need | Tool |
|---|---|
| Gradient of one scalar wrt params | .backward() / autograd.grad |
| Second derivative for meta-loss | create_graph=True double-backward |
| Full Jacobian (small dims) | jacrev/jacfwd |
| Batch of small problems | vmap |
| HVP for implicit/Hessian-based optim | jvp ∘ vjp |
| Per-sample grads (fisher, robust) | vmap over autograd.grad |
Key Takeaways¶
create_graph=Trueextends the graph for a second backward.jvp/vjpare the atoms; Jacobians and Hessians are compositions of them.- HVP ≈ 2 forwards; never materialize a full Hessian.
vmapreplaces batch loops and unlocks per-sample methods.- Verify with
torch.autograd.gradcheckon custom fns before trusting results.