Skip to content

Forward-Mode Residual Derivatives ​

TL;DR

PDE residuals differentiate the network with respect to a handful of scalar coordinates — the opposite regime from backprop. JAXPI computes them with forward-mode JVP sweeps (jaxpi.derivatives): values identical to the reverse-mode idioms to floating point rounding, residual evaluation 1.4–3× faster on multi-output models, and no tape in memory. The parameter gradient of the loss is unaffected and stays reverse-mode.

The problem: the residual stack is the inner loop ​

A Navier–Stokes-style residual needs a lot of derivatives of the network fθ:(t,x,y)↦(u,v,p) at every collocation point of every training step: the full first-derivative Jacobian plus the velocity Laplacians,

J=∂(u,v,p)∂(t,x,y)∈R3×3,uxx,uyy,vxx,vyy.

The idiomatic JAX implementation — one jacrev for the Jacobian, one jax.hessian per velocity component — does redundant work in exactly this regime:

  • jacrev runs one taped backward pass per output: forward pass, store every intermediate, transpose sweep. Three outputs, three tapes.
  • jax.hessian per component builds the full spatial Hessian, including the cross terms uxy, vxy that a Laplacian never uses — and it pays this per component, so a 3D model burns nine taped Hessian columns to keep six numbers.

None of this is wrong — the values are exact — but the shape of the computation is wrong for the shape of the problem.

Forward vs reverse: pick the mode by shape ​

The network is a composition f=fL∘⋯∘f1, so its Jacobian is the matrix product J=JL⋯J1. Autodiff never forms the factors; it applies them to vectors, in one of two orders:

JVP (forward):z˙k+1=Jkz˙kVJP (reverse):z¯k=Jk⊤z¯k+1

One JVP costs about two forward passes, stores nothing, and yields Js — the derivative of every output in one input direction s. One VJP yields w⊤J — the gradient of one output with respect to every input — but must record a tape and run a transpose sweep. A full Jacobian therefore costs nin JVPs or nout VJPs:

mapshaperight mode
loss w.r.t. parameters106→1reverse (backprop) — unchanged
residual w.r.t. coordinates3–4→2–5forward — same pass count, no tape, no transpose

Both orders evaluate the same chain-rule product, and matrix multiplication is associative — so the values are identical up to float rounding. This is a pure performance and memory decision, never an accuracy trade-off.

For second derivatives, a JVP of a JVP gives a second-order directional derivative, ∂r∂sf=r⊤(∇2f)s; seeding both levels with the same basis vector ei reads off exactly the diagonal entry ∂2f/∂ci2. Two properties make this strictly better than jax.hessian here:

  • only the diagonal is computed — no cross terms;
  • each sweep carries all output components at once, so the cost scales with the number of coordinates, not outputs. The old per-component idiom paid m× full Hessians; the forward version pays ds nested sweeps. For the 3D Taylor–Green model that turns nine taped Hessians into three tape-free sweeps.

Design decisions ​

jaxpi.derivatives deliberately stays small — three primitives that mirror how residuals are actually written:

  • value_and_jacfwd(f, argnums) traces f once with jax.linearize and evaluates one JVP per coordinate, returning values and all first derivatives — the common "call the network, then differentiate it" double evaluation disappears.
  • hessian_diag_fwd(f, argnums) is the forward-over-forward diagonal described above.
  • derivatives_fwd_1d(f, argnum, order) builds one nested tower for successive derivatives along a single coordinate.
  • Helpers take scalar coordinate arguments, matching the per-point r_net signature; batching stays outside in vmap, exactly as before.

Two variants were benchmarked and rejected:

  • Vectorizing the JVPs over the coordinate basis (a vmap over tangents instead of a Python loop) looks elegant but composes badly with the outer batch vmap: Taylor–Green regressed from 9.7 ms to 13.4 ms. XLA fuses the sequential form better; it stayed.
  • Migrating everything unconditionally. Two residual families are better off as they were, and keep their original code: first-order scalar residuals (advection, inviscid_burgers), where reverse mode needs one pass against forward's two and forward measured 0.72×; and high-order 1D chains (kdv, ks), where nested JVP towers grow like 2k and Taylor-mode jet already computes the whole derivative tower in a single pass.

During training the residual sits inside the loss, whose parameter gradient is taken by reverse mode — so the full computation is reverse-over-forward. JAX transposes JVPs cheaply; the measured step-time gains track the residual-level gains.

Results ​

Residual evaluation (vmap-ed r_net, 8192 points, one GPU), reverse-mode idiom → forward-mode helpers:

examplebeforeafterspeedup
rayleigh_taylor (4 outputs, 2D)11.1 ms3.7 ms3.0×
taylor_green (4 outputs, 3D)24.8 ms9.7 ms2.5×
lid_driven_cavity / bfs_flow3.2 ms1.5 ms2.1×
kolmogorov_flow_Re1e615.1 ms7.3 ms2.1×
taylor_green multi-stage49.3 ms24.5 ms2.0×
kolmogorov_flow6.8 ms3.7 ms1.9×
ginzburg_landau / gray_scott1.7× / 1.4×
wave / burgers / allen_cahn1.05–1.12×
kdv / ks / sod_shock_tube / advectionunchanged (already optimal)

The ranking follows the theory: gains grow with output count and spatial dimension, because those multiply the redundant reverse passes the old idiom paid.

In JAXPI ​

A 2D Navier–Stokes residual in the forward-mode idiom:

python
from jaxpi.derivatives import hessian_diag_fwd, value_and_jacfwd

def r_net(self, params, t, x, y):
    (u, v, p), (d_t, d_x, d_y) = value_and_jacfwd(
        self.neural_net, (1, 2, 3))(params, t, x, y)
    u_t, v_t, _ = d_t
    u_x, v_x, p_x = d_x
    u_y, v_y, p_y = d_y

    (u_xx, v_xx, _), (u_yy, v_yy, _) = hessian_diag_fwd(
        self.neural_net, (2, 3))(params, t, x, y)
    ...

Every example migration was verified against fixtures captured from the pre-migration code (tests/residual_equivalence.py --capture / --verify): all 17 reproduce their original residuals to ≤5×10−15 relative in float64, and tests/test_derivatives.py pins the helpers against closed-form derivatives and the reverse-mode references. The --bench mode reproduces the table above. See the jaxpi.derivatives API reference.

Released under the Apache 2.0 License.