definition theorem derivation design
The derivation behind ImplicitREDDiff, carried through. A diffusion model defines a relation as the zero set of a residual field. Inference is root-finding on that field and backpropagation is the adjoint of the root. Read through Tweedie’s formula, RED-Diff’s stop-gradient field turns out to be the exact gradient of a smoothed log-density. So the scalar energy and the multivariate energy of the original sketch are one object, and inference is a proximal step.
Implemented in
lib/VariationalDiffusion.jl/src/implicit.jlandanalytic.jl; see implicit, analytic.
Sources: original to this vault (design and analysis); builds on Song et al., Score-Based Generative Modeling through Stochastic Differential Equations, ICLR 2021, arXiv:2011.13456; Mardani, Song, Kautz & Vahdat, A Variational Perspective on Solving Inverse Problems with Diffusion Models, arXiv:2305.04391, Proposition 2; Efron, Tweedie’s Formula and Selection Bias, JASA 2011; code:
implicit.jl,analytic.jlTheory (CT-ML wiki): Bayesian Inversion · Statistical Game · Variational Free Energy · Least Fixed Point · Bisimulation
1. From functions to relations
Explicit learning fits . Implicit learning fits a relation through a residual :
The direction (what is input, what is output) is not part of the model. It is chosen at inference time, which is what makes one model serve every conditional. For a diffusion model the residual comes from the score, and the rest of this note derives it.
2. Polarity is a precision vector
Split the coordinates of into inputs , outputs and latents with
diagonal selection matrices (pairwise orthogonal), and weight
them, . In the code this is a precision vector
, which is also what precision_vector(f, polarity) returns for a
DiffusionFactor:
| coordinate is | constraint | |
|---|---|---|
| a hard input | , imposed by projection | |
| large, finite | a soft input (noisy evidence) | penalty |
| small, | an output anchored at a previous value | a proximal pull towards |
| a free output or latent | none |
Two corrections to the original sketch:
- The typical ordering is , not the reverse: inputs are pinned hardest.
- A nonzero is not a nuisance. It anchors the output at its previous value , and it is what makes inference warm-startable and branch-selecting (§6).
Write for the free coordinates and for the clamped ones.
3. The residual, and the theorem that makes it a gradient
RED-Diff’s regulariser gradient (Mardani et al., Proposition 2) is, with weight ,
the gradient of the score-matching term with the network Jacobian dropped (the “stop-gradient”). Adding the polarity term gives the residual on the free coordinates:
This is the right form of the “multivariate energy” of ImplicitREDDiff. Compared with the sketch: the data term enters as , not , because it is the gradient of ; the factor from the reparametrisation is absorbed into ; and the residual is not squared inside the expectation.
Theorem (the stop-gradient field is conservative). If is the optimal noise predictor, then
and for any fixed set of noise nodes the same identity holds node by node, up to the constant vector (the gradient of a linear function).
Proof. Tweedie’s formula for the Gaussian perturbation kernel gives . For a fixed , the chain rule gives . Hence . Averaging over nodes or integrating over commutes with , and the term does not depend on .
Consequences:
- The scalar and the vector energy are one object. Inference minimises and the residual is its gradient. is a mixture of Gaussian-smoothed negative log-densities: each noise level contributes at scale , with weight .
- RED-Diff is not “biased” in the sense the vault used to say. Its field is not the gradient of the regulariser it starts from. It is the exact gradient of , a different and well-defined prior. (RED-Diff as a Statistical Game §6 called it “approximate at its fixed point”; the precise statement is “exact for a smoothed prior”.) The Gaussian calibration of that note’s §4 is the special case where is quadratic.
- For a learned network the field need not be a gradient. is then not symmetric, does not exist, and only the residual is meaningful. So the vector formulation is the primary one: it is defined for every network, and it reduces to the scalar one exactly when the network is a true score.
The identity is checked numerically for a Gaussian mixture, whose is closed form
(analytic.jl): field and agree to , and the Jacobian is symmetric to
machine precision (test suite, “the deterministic field is the gradient of a smoothed log-density”).
4. Inference is a proximal point
With , a root of with is a stationary point of
i.e. the proximal operator of the smoothed log-prior in the metric . This makes precise the slogan of The Diffusion Family (“a statistical game whose inversion is a proximal operator”) and connects it to ProxDM and Proximal Alternatives. ProxDM learns a prox; here the prox is induced by a learned score.
5. Roots, stable roots, and the relation
also holds at maxima and saddles of . Those points are on the zero set but are not plausible configurations. The relation one wants is the set of stable roots,
and implicit_infer reports both flags, converged and stable. When the field is a
gradient, stability is the second-order condition for a local minimum.
The smoothing scale decides what relation you get
smooths the data density at scales , so the recovered relation is the ridge of a smoothed density, not the support of the data. For a thin ring of radius (mixture components of width ) smoothed at a single level, the radial mode moves inward by, to first order,
Measured on the unit circle, inferring at :
| noise levels used | largest | inferred radius |
|---|---|---|
| 0.045 | 0.997 | |
| (the package default) | 0.172 | 0.980 |
| 0.322 | 0.924 | |
| 0.584 | 0.658 | |
| (RED-Diff’s training range) | 1.000 | 0.000 |
A single noise-free level at gives and the formula predicts ; at , against .
RED-Diff’s own weighting, over , smooths a circle until its
density peaks at the centre, and then the relation is a single point. For implicit learning
the noise levels used at inference must stay below the curvature scale of the relation.
That is a design parameter (levels in field_nodes), and it is a bias–conditioning
trade-off: small means little bias but a stiff, ill-conditioned field (
scales like near the data).
6. Multi-valued relations and branches
On the circle, a hard input has two stable outputs, . Inference returns the one whose basin contains the starting point; a small anchor at a previous biases the choice. Near the two branches merge: smoothing turns the discriminant locus of Branches and the Discriminant into a region where only one root survives, and starts from either side land on it. Outside the support () inference still answers with a point near the ridge: the “closest point to the variety” behaviour that README’s table predicts for implicit learners.
Three caveats, all observed:
- With few noise samples the deterministic field breaks the symmetry, so the merged root sits slightly off the axis. That is a sample-average artefact, smaller with more nodes.
- In flat regions far from data the field is tiny, and descent can stall without converging. The solver reports this rather than hiding it.
- Training through inference (see Backpropagation through Implicit Inference §7) only moves the branches inference visits. The rest of the old relation survives as other branches.
7. The pieces, and where they are worked out
| question | note |
|---|---|
| what inference is, in all its signatures | Inference Signatures |
| the Lagrangian, adjoint and gradient, worked out | Backpropagation through Implicit Inference |
| the deterministic limit, and why it is a DEQ | Deterministic Relaxation |
| the statistical-game reading, and what is missing | The Implicit Diffusion Factor as a Statistical Game |
tab: Julia
**Docs:** [VariationalDiffusion API](https://mathstruct.org/Lenticulum.jl/dev/packages/variationaldiffusion/)
```julia
using VariationalDiffusion, LuxCore, Random, LinearAlgebra
sched = VPSDE()
θ = range(0, 2π; length = 49)[1:48]
ring = NoisePredictor(GaussianMixtureEps(sched, vcat(cos.(θ)', sin.(θ)'); s = 0.05), sched) # exact score of a circle
ps, st = LuxCore.setup(Xoshiro(0), ring)
m = ImplicitDiffusion(ring, field_nodes(Xoshiro(1), 2; samples = 8)) # levels t ∈ [0.002, 0.05]
ρ = [Inf, 0.0] # x hard input, y free output
up, _ = implicit_infer(m, [0.6, 0.5], ρ, ps, st) # start above the axis
dn, _ = implicit_infer(m, [0.6, -0.5], ρ, ps, st) # start below
(up.z[2], dn.z[2]) # ≈ (0.78, -0.78): two branches of one relation
(up.converged, up.stable) # (true, true)
```