derivation theorem algorithm

The Lagrangian of the original sketch, carried through: state equation, adjoint equation and gradient for an implicit diffusion learner. Also covered: the hard/soft clamp split, what the adjoint costs, why no noise has to be stored, and what to do when inference has not reduced the energy. Every gradient below is checked against finite differences of the full solve in lib/VariationalDiffusion.jl/test/runtests.jl.

Sources: original to this vault (design and analysis); the adjoint method as in Backpropagation by the Implicit Function Theorem and DEQ as a Relation §4; Bai, Kolter & Koltun, Deep Equilibrium Models, NeurIPS 2019; Fung et al., JFB: Jacobian-Free Backpropagation for Implicit Networks, AAAI 2022; Mardani et al. arXiv:2305.04391; code: implicit.jl (implicit_pullback), analytic.jl (epsilon_jacobian, epsilon_vjp_params)

Theory (CT-ML wiki): Lens · Reverse Derivative Category · Statistical Game · Compiler Correctness

1. Setup

From Implicit Diffusion Learners: free coordinates (finite precision), clamped coordinates (precision ), and the residual on

with the noise nodes fixed (sample-average approximation; why this matters is §5). Inference returns with and . Its Jacobian on the free block is

A downstream loss (e.g. ) has cotangent .

2. The Lagrangian

State equation:

This is the inference problem: solve it first, with any of the point signatures of Inference Signatures.

Adjoint equation:

One linear solve with the transpose of the free-block Jacobian. For an exact score is symmetric (Implicit Diffusion Learners §3), so the transpose is irrelevant; for a learned network it is not symmetric, and the transpose matters. Solve it densely for small (form from VJPs) or matrix-free with GMRES for large ; each GMRES iteration costs one VJP of per node.

Gradient: read everything off

By the implicit function theorem, along every input direction. Embedding in with zeros on :

w.r.t.gradientneeds
network parameters one parameter-VJP per node
hard-clamped inputs the off-diagonal Jacobian block
soft targets / anchors nothing extra
precisions nothing extra

The last row means the polarity weights are learnable. Soft-clamp precisions can be trained by the same adjoint, so “how much to trust an input” becomes a parameter rather than a hyperparameter.

implicit_pullback(m, sol, z₀, ρ, z̄, ps, st) returns (z₀ = …, ρ = …, ps = …). In the backpropagation signature of the sketch it is

and the incoming energy does not need its own cotangent (§6).

3. Why this is the right derivative

The adjoint gives the exact derivative of the map defined by the deterministic residual, wherever is invertible. Checked against central finite differences of the full solve (tolerance ), on a circle relation with exact score:

derivativeadjointfinite difference
(hard input)−0.06184743−0.06184742
(a mixture mean, i.e. )+0.02653255+0.02653255
(soft input)+0.02502245+0.02502245
+0.01145560+0.01145560
−0.00007650−0.00007650

For a Gaussian prior everything is linear: and , matched to .

4. The cost

Per backward pass: one Jacobian evaluation (or one GMRES solve), and one parameter-VJP of per noise node. No unrolling, so memory is independent of the number of inference iterations, the defining advantage of implicit differentiation. With nodes this is network VJPs, the same as one step of RED-Diff with samples, done once rather than once per iteration.

5. Does the noise from inference have to be stored?

No, and the reason is structural. There are three cases:

  1. Fixed nodes (what the package does). The noise is part of the model, a fixed quadrature of the expectation. The backward pass re-evaluates the Jacobians at with the same nodes; nothing from the forward trajectory is needed except .
  2. Fresh noise in inference (RED-Diff as published). The forward pass is a stochastic approximation whose limit is a root of the expected field . The IFT applies to , at the limit point. The trajectory’s noise is irrelevant. The backward pass needs unbiased estimates of and , which fresh noise provides.
  3. The catch in case 2. involves an inverse of an expectation, and . Plugging a noisy Jacobian into the solve gives a biased gradient. Fix: use many nodes in the backward pass, or the same fixed node set for the Jacobian and the residual. Fixing the nodes is the simplest way to make the whole thing exact.

So save , not the noise. If the noise is fixed it is part of ‘s computational graph anyway.

6. When the energy is not reduced

The IFT is a statement about a root. At a point that is not a root it computes the derivative of nothing in particular. implicit_pullback therefore refuses non-converged solutions. The options, from most to least principled:

situationwhat to do
converged, stablethe adjoint (§2) is exact
converged, unstable (saddle or maximum of )do not backpropagate; the point is not in the relation. Restart inference, or anchor with
not converged, residual smallcontinue iterating from the current state (warm start), then backpropagate
not converged, budget exhaustedbackpropagate through the unrolled last steps (exact for the truncated map), or use Jacobian-free backpropagation (replace by the identity, as in JFB) as a descent direction
no root at all (input outside the relation)inference returns the nearest ridge point; the residual reports the distance

The loss-side question (“what if inference increased the energy?”) does not arise for the solver here: Newton steps are only taken where , with backtracking on . Elsewhere the step is a descent step on when the field is a gradient. A learned non-conservative field has no energy to decrease, only a residual. That is why convergence is defined on and not on .

Observed in training (§7): the first prototype dropped about 60% of solves as non-converged, and its gradients blew the model up until those were excluded. With the stability guard and the Newton/descent switch, the package’s solver dropped none.

7. Learning a relation by backpropagating through inference

The experiment that closes the loop. Start from the circle, a ring of 64 mixture components whose means are . Pairs come from a different curve, . Train by Adam on , where is the inferred output, with gradients from the adjoint only and no denoising loss at all.

beforeafter 150 epochs
training loss (mean squared error)
test RMSE, warm start0.590.012
test RMSE, cold start 0.590.26
dropped solves—0

The cold-start number is one test point () landing on the upper arc of the old circle, which training never touched. Bilevel training reshapes the branch that inference visits and leaves the others. An implicit learner trained this way is a relation with several branches, and inference returns the branch of its starting basin. To fit a whole relation, either train generatively as well (the statistical-game semantics of The Implicit Diffusion Factor as a Statistical Game), or anchor the output () at an amortised guess.

8. Where AD is needed, and where it is not

quantityexact score (closed form)learned network
forward passforward pass
(inference Newton steps, adjoint solve)closed forminput JVP/VJP, or finite differences for small
(parameter gradient)closed formreverse-mode AD: one VJP per node

VariationalDiffusion has no AD dependency of its own: epsilon_jacobian has a finite-difference fallback, and epsilon_vjp_params is an interface method that closed-form predictors implement exactly. A Lux network gets both from the backend named in its predictor, NoisePredictor(…; ad = AutoZygote()) (or Enzyme, ForwardDiff, Mooncake, Reactant), through a package extension (backends). A trained 5k-parameter MLP on the circle reproduces the finite-difference derivatives of this section to 8 digits (examples/circle_mlp.jl). Compare RED-Diff as a Statistical Game §3, where the stop-gradient was “the difference between a package with an AD dependency and one without”. Training through inference brings back exactly one VJP per node, no more.

Related: Implicit Diffusion Learners, Inference Signatures, Deterministic Relaxation, The Implicit Diffusion Factor as a Statistical Game, DEQ as a Relation, Backpropagation by the Implicit Function Theorem, implicit, backends