VariationalDiffusion

A diffusion model as a factor: wrap a trained noise predictor and use it as a prior, with inference by RED-Diff's proximal solve.

This package consumes a trained $\varepsilon_\theta$; it does not train one. There is no automatic-differentiation dependency, because RED-Diff's stop-gradient means the network is only ever evaluated forwards.

The pieces

sched = VPSDE()                                   # the forward process
pred  = NoisePredictor(my_unet, sched)            # ε_θ(x, t) — any Lux model
prox  = REDDiff(λ = 0.25, steps = 200)            # the inversion
f     = DiffusionFactor((obs = 4, hidden = 4), pred; prox)

The schedule

VPSDE is the variance-preserving SDE of Song et al.: $x_t = \alpha_t x_0 + \sigma_t \varepsilon$ with $\alpha_t^2 + \sigma_t^2 = 1$. Query it with alpha, sigma, snr, perturb, sample_time.

The predictor

NoisePredictor wraps any LuxCore.AbstractLuxLayer and derives two things from it:

functionisnote
epsilonthe network itselfthe only place it is evaluated
score$-\varepsilon_\theta/\sigma_t$$\approx \nabla_x \log p_t(x)$
denoise$(x - \sigma_t\varepsilon_\theta)/\alpha_t$Tweedie — the MMSE denoiser

Parameters are the wrapped model's, untouched: LuxCore.setup(rng, pred) == LuxCore.setup(rng, unet).

The factor

DiffusionFactor carves one vector space into named channel blocks. The Polarity decides which blocks are clamped and how hard, which is exactly the selection-matrix formulation $P = \rho_{in}P_{in} + \rho_{out}P_{out} + \rho_{latent}P_{latent}$ — inpainting, in other words. Observed() defaults to $\rho = \infty$, a hard clamp applied by projection.

The inversion runs REDDiff and returns a DiracBelief, because RED-Diff's variational family is a point mass. That is the paper's choice, not a simplification here.

Choosing λ

λ is presented in the paper as a tuned hyperparameter (they use 0.25). For a Gaussian data distribution it is derivable — there is exactly one value making the implied prior correct, and calibrate_lambda(schedule, v₀) returns it.

Worth knowing what tuning λ actually does: it sets how strong the learned prior is. A single scalar can only calibrate one eigendirection of a correlated prior, so it over-regularises high-variance directions and under-regularises low-variance ones.

Small networks, any AD backend

A diffusion model here is a prior over the joint space of one factor: a few coordinates, not an image. It is evaluated many times per query, so small models are the point, and the architecture is free. An MLP with a sinusoidal time embedding is the tested baseline; input adapts $(x, t)$ to whatever the model expects.

The implicit learner needs two derivatives of the network: the input Jacobian (Newton steps, adjoint solve) and a parameter VJP (backward pass). The backend is a field of the predictor:

using DifferentiationInterface, Zygote          # or Enzyme, ForwardDiff, Mooncake, …
pred = NoisePredictor(mlp, VPSDE(); input, ad = AutoZygote())

using Reactant                                  # XLA-compiled forward and Enzyme VJPs
pred = NoisePredictor(mlp, VPSDE(); input, ad = AutoReactant())

Both are package extensions; without them the Jacobian falls back to finite differences. examples/circle_mlp.jl trains a 5k-parameter MLP on a circle in about 20 seconds on a CPU, then infers both branches and checks the adjoint against finite differences. For Enzyme on Lux use AutoEnzyme(; mode = Enzyme.set_runtime_activity(Enzyme.Reverse)).

Energy-parametrised models

NoisePredictor(EnergyNetwork(net), VPSDE(); input, ad = AutoZygote()): the network outputs a scalar energy and $\varepsilon_\theta = \sigma_t\nabla_x E_\theta$, so the score is conservative by construction. The implicit learner's field is then the gradient of implicit_energy, a real energy for the learned relation. Train with denoising_gradient (the loss gradient as one mixed second derivative, no nested AD in the loop) and any optimiser.

Proximal diffusion models

ProxDM (Fang et al. 2025) queries the prior through $\operatorname{prox}_{-\lambda\log p_t}$ instead of its score. proximal is the interface, ProxNetwork a learned prox ($v - \sqrt\lambda\,\varepsilon_\theta(v; t, \lambda)$), MixtureProx the exact one for a Gaussian mixture. proxdm_sample is the paper's sampler (PDA and PDA-hybrid), prox_infer the implicit learner's query by half-quadratic splitting, proximal_matching_loss the training loss.

Known gaps

  • prox_infer has no adjoint yet, and ProxDM is not yet a DiffusionFactor inversion.
  • The inversion returns a point, so uncertainty does not propagate.
  • The message is a posterior rather than a likelihood, so it double-counts on a variable of degree greater than one.

API

VariationalDiffusion.VariationalDiffusion — Module
VariationalDiffusion

A diffusion model as a statistical game: the third of the three Implicit Learners.md families, and the only one whose Bayesian inversion is neither exact nor a root-find but a proximal solve.

The chain

  1. schedule.jl — the VP-SDE of Song et al. 2021 (arXiv:2011.13456), presented through its perturbation kernel $p_{0t}(x_t\mid x_0) = \mathcal{N}(\alpha_t x_0, \sigma_t^2 I)$ with $\alpha_t^2 + \sigma_t^2 = 1$.

  2. predictor.jl — NoisePredictor wraps any Lux model as $\varepsilon_\theta(x,t)$ and derives from it the score $-\varepsilon_\theta/\sigma_t$ and Tweedie's denoiser $(x - \sigma_t\varepsilon_\theta)/\alpha_t$.

  3. reddiff.jl — REDDiff, the proximal operator of Mardani et al. 2023 (arXiv:2305.04391): variational inference with a point-mass posterior, whose regulariser gradient is $\mathbb{E}_{t,\varepsilon}[\lambda_t(\varepsilon_\theta(x_t,t)-\varepsilon)]$ under a stop-gradient.

  4. factor.jl — DiffusionFactor, where ImplicitREDDiff.md's selection matrices $P_{in} + P_{out} + P_{latent} = \mathrm{Id}$ become a LenticulumCore.Polarity and

    \[E(x_0, x) = \underbrace{\mathbb{E}_{t,\varepsilon}\bigl[\omega(t)\|\varepsilon_\theta(\alpha_t x + \sigma_t\varepsilon,t)-\varepsilon\|^2\bigr]}_{\text{entropy } \mathbf{H}^c} \;+\; \underbrace{\tfrac12\|P(x_0-x)\|^2}_{\text{energy } \mathbf{l}^c}\]

Two things worth knowing before reading the code

No automatic differentiation anywhere. RED-Diff's stop-gradient means the denoiser's Jacobian is never formed, and the clamp's gradient is $P^2(x-x_0)$ in closed form. So the entire inversion is forward passes plus arithmetic, and this package's dependencies are LuxCore, Random and LinearAlgebra. Lux itself is not a dependency — any Lux model is an AbstractLuxLayer, which is all the wrapper needs.

The inversion returns a DiracBelief, because RED-Diff's variational family is $\mathcal{N}(\mu,\sigma^2 I)$ with $\sigma\to 0$. That is the paper's own choice, not a simplification made here, and it has consequences on a graph — see factor.md §5.

Concept notes are in vault/Families/Diffusion/; per-file implementation notes sit next to each source file, per vault/Start Here.md.

source
VariationalDiffusion.AbstractNoiseSchedule — Type
abstract type AbstractNoiseSchedule

A forward corruption process, presented through its perturbation kernel rather than its SDE: everything downstream needs only alpha(s, t) and sigma(s, t).

That is a deliberate narrowing. The SDE is the definition; the Gaussian marginal is what RED-Diff and score matching actually touch, and a schedule whose marginal is not Gaussian would not fit this interface. drift and diffusion are provided for completeness and are not used by the inversion.

source
VariationalDiffusion.DiffusionFactor — Type
DiffusionFactor(blocks::NamedTuple, predictor; prox = REDDiff())

A factor whose prior over its own state space is a diffusion model, and whose Bayesian inversion is a RED-Diff proximal solve.

blocks names the channels and their dimensions; they tile one vector space $\mathbb{R}^n$, $n = \sum_i \dim_i$, in declaration order. predictor is a NoisePredictor over that whole $\mathbb{R}^n$ — one network for the joint state, not one per channel, which is what makes this a relation rather than a collection of conditionals.

f = DiffusionFactor((obs = 4, hidden = 4), NoisePredictor(unet, VPSDE());
                    prox = REDDiff(λ = 1.0, steps = 300))

This is the third of the three [Implicit Learners] families — the diffusion one — and the only one whose inversion is neither exact nor a root-find. LenticulumCore anticipated it: LenticulumCore.ProximalInversion exists in lens.jl and names this package.

source
VariationalDiffusion.DiffusionModel — Type
DiffusionModel(factor, polarity)

The open model the polarity selects: an implicit relation on $\mathbb{R}^n$, restricted to predicting the unobserved block from the observed ones.

Not pure. The unobserved and latent blocks are both reconstructed by the prox, and the latent ones are exactly AutoBayes' $\llbracket c \rrbracket$ — internal coordinates the inversion has to fill in and nobody reads.

source
VariationalDiffusion.EnergyNetwork — Type
EnergyNetwork(model)

Marks model as a scalar energy $E_\theta(x, t)$ rather than a noise predictor. Wrapped in a NoisePredictor, it defines

\[\varepsilon_\theta(x, t) = \sigma_t\,\nabla_x E_\theta(x, t), \qquad s_\theta(x, t) = -\nabla_x E_\theta(x, t),\]

so the score is conservative by construction:

pred = NoisePredictor(EnergyNetwork(net), VPSDE(); input, ad = AutoZygote())

net sees input(x, t) and returns one value per column of x (a 1 × B output, or a scalar for a single point). Parameters are the network's own. Every function of the package that takes a NoisePredictor accepts this one; the implicit learner additionally gains a scalar energy, implicit_energy.

Derivatives need an AD backend in ad, and second derivatives at that: the input Jacobian is a Hessian and the parameter VJP a mixed derivative. AutoZygote() and DifferentiationInterface.SecondOrder(AutoForwardDiff(), AutoZygote()) work with Lux layers (load DifferentiationInterface and the backend); Enzyme's forward-over-reverse does not yet.

source
VariationalDiffusion.FieldNodes — Type
FieldNodes(t, ε, w)

The fixed quadrature nodes $(t_k, \varepsilon_k, w_k)$ that turn RED-Diff's stochastic gradient into a deterministic vector field. Fixing them is sample-average approximation: the field, its root and its Jacobian become ordinary deterministic objects, which is what the implicit function theorem needs. Build them with field_nodes or noisefree_nodes.

source
VariationalDiffusion.GaussianMixtureEps — Type
GaussianMixtureEps(schedule, μ₀::AbstractMatrix; s = 0.05)

The exact noise predictor $\varepsilon^\ast$ of the data distribution $\frac1J\sum_j\mathcal N(\mu_j, s^2 I)$ under schedule. Columns of μ₀ (an $n\times J$ matrix) are the initial component means; they are the layer's parameters (ps.μ), so a mixture can be moved by gradient descent. The component width s is fixed.

Wrap it in a NoisePredictor with the same schedule. Its input Jacobian (epsilon_jacobian) and parameter VJP (epsilon_vjp_params) are exact.

source
VariationalDiffusion.ImplicitDiffusion — Type
ImplicitDiffusion(predictor, nodes; λ = 1.0)

The relation $R_\theta = \{z : r(z) = 0\}$ defined by a noise predictor and a fixed node set, with RED-Diff's weighting $\lambda_t = \lambda\sigma_t/\alpha_t$.

source
VariationalDiffusion.ImplicitProx — Type
ImplicitProx(nodes; λ = 1.0, tol = 1e-9, maxiters = 200, step = 0.05)

The deterministic implicit solver as a DiffusionFactor's inversion, in place of REDDiff:

f = DiffusionFactor((x = 1, y = 1), pred; prox = ImplicitProx(field_nodes(rng, 2)))

Inversion is then implicit_infer on the factor's state space with the polarity's precisions (Inf = hard clamp). Compared with RED-Diff it is deterministic, it reports convergence and stability (implicit_solution), its free energy is deterministic, and it has a backward pass (implicit_factor_pullback). nodes must have the factor's state dimension.

source
VariationalDiffusion.ImplicitSolution — Type
ImplicitSolution

What inference returns: the state z, the residual norm on the free coordinates, the number of iterations, and two flags.

  • converged — the residual is below tolerance. If not, z is still returned (an anytime answer) but the implicit function theorem does not apply at it.
  • stable — the symmetric part of $J_{FF}$ is positive definite: a strict local minimum of the energy when the field is a gradient. A root that is not stable is a saddle or a maximum of the energy — on the zero set of $r$, but not a member of the relation one wants.
source
VariationalDiffusion.MixtureProx — Type
MixtureProx(mixture::GaussianMixtureEps; tol = 1e-10, maxiters = 100)

The exact proximal operator of $-\log p_t$ for the Gaussian-mixture data distribution of GaussianMixtureEps, by a Newton solve with backtracking on the prox objective. It plays the role a perfectly trained ProxNetwork would, so samplers and proximal inference can be checked against arithmetic. Uses $\nabla\log p_t$ directly (not $\varepsilon/\sigma_t$), so it is valid at $t = 0$. Where $-\log p_t$ is not convex the prox can be multivalued; the solve returns the local minimiser reached from $v$.

source
VariationalDiffusion.NoisePredictor — Type
NoisePredictor(model, schedule; input = default_input)

$\varepsilon_\theta(x,t)$ — a wrapper making any Lux model into a diffusion prior.

model is any LuxCore.AbstractLuxLayer; parameters and state are delegated to it untouched, via AbstractLuxWrapperLayer, so ps for a NoisePredictor is ps for the model. Wrapping adds no parameters of its own.

input adapts the (x, t) pair to whatever the model expects, defaulting to the tuple (x, t). A model taking a concatenated time channel, a sinusoidal embedding or a NamedTuple is accommodated by passing a different input; the wrapper deliberately knows nothing about time embeddings, because that is the model's business.

pred = NoisePredictor(my_unet, VPSDE())
ps, st = LuxCore.setup(rng, pred)          # == LuxCore.setup(rng, my_unet)
ε, st  = epsilon(pred, x, 0.3, ps, st)
LuxCore, not Lux

This package depends on LuxCore, the interface package, exactly as Lenticulum does. Any Lux model is an AbstractLuxLayer, so Lux itself is never needed here — see predictor.md §1.

source
VariationalDiffusion.ProxNetwork — Type
ProxNetwork(model, schedule; input = (v, t, λ) -> (v, t, λ))

A learned proximal operator in ProxDM's parametrisation $f_\theta(v; t, \lambda) = v - \sqrt\lambda\,\varepsilon_\theta(v; t, \lambda)$: the network predicts the normalised residual of the prox, conditioned on two scalars, $t$ and $\lambda$. model is any Lux layer; input adapts $(v, t, \lambda)$ to what it expects. Parameters are the model's. Train it with proximal_matching_loss.

source
VariationalDiffusion.REDDiff — Type
REDDiff(; λ = 0.25, steps = 200, lr = 0.1, samples = 1, adam = true,
          β₁ = 0.9, β₂ = 0.999, ϵ = 1e-8, rng = Random.default_rng())

The RED-Diff proximal operator: configuration for the inner optimisation that is the Bayesian inversion.

fieldmeaning
λregulariser strength; enters as $\lambda_t = \lambda\sigma_t/\alpha_t$
stepsinner iterations L of Algorithm 1
lrstep size
samplesMonte-Carlo draws of $(t,\varepsilon)$ per step (the paper uses 1)
adamAdam vs plain gradient descent; the paper uses Adam

λ = 0.25 is the value Mardani et al. tuned across their experiments. It is not a universal constant, and for a Gaussian data distribution there is exactly one λ that makes the implied prior correct — see calibrate_lambda and reddiff.md §4, which is the sharpest thing this package has to say about RED-Diff.

source
VariationalDiffusion.VPSDE — Type
VPSDE(; βmin = 0.1, βmax = 20.0, tmin = 1e-3)

The variance-preserving SDE. Defaults are Song et al.'s, which are in turn DDPM's discretisation in the continuous limit.

tmin is a sampling floor, not part of the mathematics: at t = 0 we have σ_t = 0, the score -ε_θ/σ_t is a division by zero, and the RED-Diff weight λ_t = λσ_t/α_t degenerates. Every published implementation carries such a floor and most do not say so. Times are drawn from [tmin, 1]; see sample_time.

source
LenticulumCore.energy — Method
LenticulumCore.energy(f::DiffusionFactor, x, a, y, ps, st)

The graded energy at a point: (clamp = ½‖P(x₀-x)‖², score = one MC sample of the score-matching loss).

x is the full state vector, y the reference $x_0$, and a the polarity (the slot AutoBayes reserves for $\llbracket c \rrbracket$, used here to carry the thing that determines $P$). That reuse is ugly and is discussed in factor.md §4.

source
LenticulumCore.energyspace — Method
LenticulumCore.energyspace(::DiffusionFactor)

$E_c = \mathbb{R}^2$, graded as (clamp, score):

summandisrole
clamp$\tfrac12|P(x_0-x)|^2$the energy $\mathbf{l}^c$ — data consistency
score$\mathbb{E}_{t,\varepsilon}[\omega(t)|\varepsilon_\theta - \varepsilon|^2]$the entropy $\mathbf{H}^c$ — the learned prior

The split is not cosmetic and it is the reading Implicit Learners.md §"Diffusion" already gives: the clamp term is pointwise and depends on the data, so it is an energy; the score-matching term depends on the learned distribution rather than on the data point, so it is an entropy. Getting this backwards would put the prior in the energy and break the counting correction of Bethe Free Energy.md.

source
LenticulumCore.invert — Method
LenticulumCore.invert(lens, π, inputs, ps, st) -> (DiracBelief, st)

Run the RED-Diff prox and return the reconstructed unobserved block.

The return type is a DiracBelief, and that is faithful rather than lazy: RED-Diff's variational family is $q = \mathcal{N}(\mu, \sigma^2 I)$ with $\sigma \to 0$, so the posterior it computes is a point mass. The consequence is recorded in factor.md §5 — a Dirac message dominates every combine it meets, so a diffusion factor in a graph overrides its neighbours rather than negotiating with them.

source
LenticulumCore.supported_polarities — Method
LenticulumCore.supported_polarities(f::DiffusionFactor)

The n polarities with one channel Unobserved() and the rest Observed().

This under-reports what the factor can do. LenticulumCore.supports_polarity accepts any assignment with at least one unobserved channel, including Latent() ones, because a diffusion prior over the joint space can inpaint any subset from any other subset — that is the entire appeal of using one. The enumeration is truncated because a scheduler needs a listable set and the full set has $3^n - 2^n$ elements. See factor.md §3.

source
Mycelium.factor_message — Method
Mycelium.factor_message(f::DiffusionFactor, target, polarity, inputs, prior, ps, st)

The factor → variable message: a DiracBelief on target, from a RED-Diff prox over the joint space.

[!warning] This is the posterior, not the likelihood Every other factor in this project returns a likelihood here, with the prior divided out, because a variable of degree d would otherwise count the prior d times (gaussian.jl's invert docstring). A diffusion factor cannot divide its prior out — the prior is a neural network and there is no subtraction available in canonical form.

So on a graph where the target variable has degree > 1, this message double-counts. It is correct for a degree-1 target (the usual inverse-problem setting: one prior, one measurement) and approximate otherwise. factor.md §5.

source
Mycelium.local_free_energy — Method
Mycelium.local_free_energy(f::DiffusionFactor, msgs, ps, st)

The factor's Bethe contribution, evaluated at the prox's own solution — i.e. at the point the inversion actually returned, which is the only point where the two summands are comparable.

Because the variational posterior is a Dirac, its entropy is $-\infty$ and the $-H(b_c)$ term of the Bethe free energy is not the posterior entropy: the score summand plays that role instead, per Implicit Learners.md. This is the one factor in the project whose free energy is not a closed form, and it is Monte-Carlo noisy by construction.

source
VariationalDiffusion.assemble_state — Method
assemble_state(f, inputs, π) -> x₀

Build the reference configuration $x_0$ of $\tfrac12\|P(x_0-x)\|^2$ by laying each channel's incoming belief into its block.

A DiracBelief contributes its value; anything with a mean contributes that; a channel with no message contributes zeros. The prior π fills the target block when it carries a point. Channels whose precision is zero never read their entry, so the zeros are not a silent default — they are multiplied out.

source
VariationalDiffusion.blockranges — Method
blockranges(f::DiffusionFactor) -> NamedTuple

The index range of each channel inside $\mathbb{R}^n$. These are the selection matrices $P_{in}, P_{out}, P_{latent}$ of ImplicitREDDiff.md, stored as ranges because a diagonal 0/1 matrix is a wasteful way to write a range.

source
VariationalDiffusion.calibrate_lambda — Function
calibrate_lambda(s::AbstractNoiseSchedule, v₀ = 1.0; nodes = 2000) -> Real

The unique λ for which RED-Diff's regulariser is the exact negative log-prior gradient of a Gaussian data distribution $\mathcal{N}(0, v_0 I)$.

For that data distribution $\varepsilon_\theta$ is available in closed form, and the regulariser gradient collapses to a linear shrinkage:

\[\nabla_x\mathrm{reg}(x) \;=\; \kappa\,x, \qquad \kappa \;=\; \lambda\int_{t_{\min}}^{1}\frac{\sigma_t^2}{\alpha_t^2 v_0 + \sigma_t^2}\,dt\]

whereas the true prior gradient is $x/v_0$. Setting $\kappa = 1/v_0$ gives the returned value. Computed by the trapezoidal rule; the integrand is smooth on $[t_{\min},1]$.

[!important] Why this matters beyond the test suite λ is presented in the paper as a tuned hyperparameter, and here it is derived. That tells you what tuning λ is really doing: choosing how strong the learned prior is, with a wrong value biasing every posterior by a known factor. See reddiff.md §4.

source
VariationalDiffusion.denoise — Method
denoise(p::NoisePredictor, x, t, ps, st) -> (x̂₀, st)

Tweedie's formula: $\hat x_0 = (x - \sigma_t\varepsilon_\theta(x,t))/\alpha_t$, the posterior mean $\mathbb{E}[x_0 \mid x_t = x]$.

[!note] This is exact, not heuristic For a Gaussian data distribution the test suite checks denoise against the closed-form linear-Gaussian posterior mean and they agree to floating point. Tweedie's formula turns a noise predictor into an MMSE denoiser, and that is the sense in which a diffusion model is a prior: it is the object the RED-Diff regulariser scores against.

source
VariationalDiffusion.denoising_gradient — Method
denoising_gradient(p::NoisePredictor, x₀, t, ε, ps, st) -> (loss, ps̄, st)

The denoising loss $\frac1B\lVert\varepsilon_\theta(\alpha_t x_0 + \sigma_t\varepsilon, t) - \varepsilon\rVert^2$ over a batch (columns of x₀, t a scalar or a 1 × B row) and its gradient with respect to the parameters, through epsilon_vjp_params. For an EnergyNetwork this is the training step: the loss already contains a derivative of the network, and this computes the loss gradient as one mixed second derivative, with no nested AD in the training loop. Apply ps̄ with any optimiser.

source
VariationalDiffusion.denoising_loss — Method
denoising_loss(p, x₀, t, ε, ps, st) -> (ℝ, st)

$\|\varepsilon_\theta(\alpha_t x_0 + \sigma_t\varepsilon, t) - \varepsilon\|^2$, one Monte-Carlo sample of the score-matching objective at a single (t, ε).

Unweighted: the weighting $\omega(t)$ belongs to whoever is taking the expectation, and different weightings mean different things — the ELBO, the perceptual-quality objective, and RED-Diff's regulariser are the same integrand with three different $\omega$. See schedule.md §3.

source
VariationalDiffusion.density_lambda — Method
density_lambda(schedule, nodes::FieldNodes) -> λ

The weighting $\lambda = 1 / \sum_k w_k\,\sigma_k^2/\alpha_k^2$ under which the implicit learner's prior term is a weighted average of smoothed negative log-densities, $U(z) \approx \sum_k \bar w_k\,(-\log p_{t_k}(\alpha_k z + \sigma_k\varepsilon_k))$ with $\sum_k \bar w_k = 1$: a proper negative log-density, in the same units as the clamp's Gaussian log-likelihood. With it, implicit_laplace returns covariances in data units. For queries whose coordinates are all hard inputs or free outputs ($\rho_i \in \{0, \infty\}$), λ does not change the answers, only the curvature; with soft evidence it sets the balance between prior and evidence.

source
VariationalDiffusion.drift — Method
drift(s::VPSDE, x, t)

The SDE's drift coefficient $f(x,t) = -\tfrac12\beta(t)x$; see diffusion for $g(t)$.

Provided for completeness and to document what the schedule is. Nothing in this package integrates the SDE — RED-Diff replaces sampling with optimisation, which is the whole point (RED-Diff as a Statistical Game.md §2).

source
VariationalDiffusion.energy — Method
energy(p::NoisePredictor{<:EnergyNetwork}, x, t, ps, st) -> (E, st)

The network's energy $E_\theta(x, t)$: a scalar for a single point, the sum over columns for a batch. One forward pass; no derivative needed.

source
VariationalDiffusion.epsilon — Method
epsilon(p::NoisePredictor, x, t, ps, st) -> (ε̂, st)

One forward pass of the wrapped network: $\varepsilon_\theta(x, t)$.

This is the only place the network is evaluated forwards. RED-Diff's stop-gradient means sampling and RED-Diff need no reverse pass through model. Only the implicit learner's backward pass does (a parameter VJP and an input Jacobian), and it goes through the predictor's ad backend, provided by a package extension (predictor.md §5).

source
VariationalDiffusion.epsilon_jacobian — Method
epsilon_jacobian(pred::NoisePredictor, x, t, ps, st) -> Matrix

$\partial\varepsilon_\theta/\partial x$ at $(x, t)$. The generic method uses central finite differences ($O(n)$ forward passes; fine for small $n$), or the predictor's AD backend when it has one (ad = …); closed-form predictors override it. For an exact score the matrix is symmetric — it is $-\sigma_t$ times a Hessian of $\log p_t$ — and the backward pass does not assume that, because a learned network's is not.

source
VariationalDiffusion.epsilon_vjp_params — Method
epsilon_vjp_params(pred::NoisePredictor, x, t, ps, st, w) -> ps̄

The vector–Jacobian product $w^\top\,\partial\varepsilon_\theta(x,t)/\partial\theta$, shaped like ps. This is the one place the backward pass needs a derivative with respect to the network's parameters.

  • Closed-form predictors define it directly (e.g. GaussianMixtureEps).
  • Any other model gets it from an AD backend chosen by the user: build the predictor with NoisePredictor(model, schedule; ad = AutoZygote()) (or AutoEnzyme(), AutoMooncake(), AutoForwardDiff(), … — any ADTypes backend) and load DifferentiationInterface together with the backend package. This package itself depends on none of them.

Without either, it throws.

source
VariationalDiffusion.field_nodes — Method
field_nodes(rng, n; levels = range(0.002, 0.05; length = 8), samples = 4, antithetic = true)

samples noise draws at each noise level in levels, equally weighted. With antithetic each draw $\varepsilon$ is paired with $-\varepsilon$: the control-variate term $-\sum w_k\lambda_k\varepsilon_k$ of the field then cancels exactly, so the field has no constant tilt.

The levels matter more than anything else here. A relation is only recovered at noise scales below its own curvature scale: the default stops at $t = 0.05$ ($\sigma_t\approx0.16$). RED-Diff's training-time range $t\in[t_{\min},1]$ smooths a unit circle into a blob whose density peaks at the centre (Implicit Diffusion Learners.md §5).

source
VariationalDiffusion.implicit_energy — Method
implicit_energy(m::ImplicitDiffusion, z, ps, st) -> (U, st)

The scalar whose gradient is the implicit learner's field, for an energy-parametrised predictor:

\[U(z) = \sum_k w_k\,\lambda_{t_k}\Bigl[\tfrac{\sigma_k}{\alpha_k}\,E_\theta(\alpha_k z + \sigma_k\varepsilon_k,\ t_k) - \varepsilon_k^\top z\Bigr], \qquad \nabla U = g = \texttt{prior\_field}(m, z).\]

So the learned relation has an energy, and a query has a loss: $U(z) + \tfrac12\lVert P(z - z_0)\rVert^2$ on the free coordinates. For a noise predictor that outputs ε directly no such $U$ exists; see energy.md §3.

source
VariationalDiffusion.implicit_factor_pullback — Method
implicit_factor_pullback(f, polarity, inputs, π, z̄, ps, st)
    -> (inputs = NamedTuple, ps = ps̄, precisions = NamedTuple)

The backward pass of the factor's inversion, by the adjoint of implicit_pullback. z̄ is the cotangent of a downstream loss with respect to the inferred state, given per channel as a NamedTuple (e.g. (y = ȳ,); channels left out have zero cotangent).

Returns the cotangents with respect to

  • each channel's incoming point (inputs): the clamped value for an observed channel, the anchor for a soft one;
  • the predictor's parameters (ps);
  • each channel's precision (precisions; zero for hard-clamped channels, whose precision is infinite and not a differentiable quantity).

Throws if the solve did not converge, as implicit_pullback does.

source
VariationalDiffusion.implicit_infer — Method
implicit_infer(m, z₀, ρ, ps, st; z_init = z₀, tol = 1e-9, maxiters = 200, step = 0.05)
    -> (ImplicitSolution, st)

Find a stable root of the residual on the free coordinates, warm-started from z_init.

Each iteration takes a Newton step $-J_{FF}^{-1}r_F$ with backtracking on $\|r\|$ when the symmetric part of $J_{FF}$ is positive definite, and otherwise a descent step $-\texttt{step}\cdot r_F$ along the field (energy descent when the field is a gradient). Newton is only trusted in locally convex regions so that it is not drawn to saddles and maxima.

z₀ holds the clamp values on hard coordinates and the soft targets elsewhere; ρ is the precision vector, Inf for a hard clamp (cf. precision_vector).

source
VariationalDiffusion.implicit_laplace — Method
implicit_laplace(m, sol, ρ, ps, st) -> (mean, cov, free)

The Laplace approximation of the query at the answer sol: mean sol.z, and on the free coordinates the covariance

\[\Sigma_{FF} = \bigl(\operatorname{sym} J_{FF}\bigr)^{-1}, \qquad J_{FF} = \partial_{z_F} r_F(z^\star) = \partial_{z_F} g_F + \operatorname{diag}(\rho_F^2),\]

the inverse curvature of the query's energy (for an energy-parametrised predictor $J_{FF}$ is already symmetric; otherwise its symmetric part is used). Hard-clamped coordinates have zero variance. free marks the free coordinates.

Units. The energy is the clamp $\tfrac12\lVert P(z - z_0)\rVert^2$ (a Gaussian log-likelihood with precision $\rho^2$) plus the prior term, a weighted sum of smoothed negative log-densities scaled by m.λ. cov is a calibrated posterior covariance when that prior term is a calibrated negative log-prior. Build the model with λ = density_lambda(schedule, nodes) for that: the covariance is then calibrated to the density smoothed at the field's noise levels (wider than the data distribution by those levels). With another λ it is a relative measure of how sharply the answer is determined. Throws unless sol converged to a stable answer.

source
VariationalDiffusion.implicit_pullback — Method
implicit_pullback(m, sol, z₀, ρ, z̄, ps, st) -> (z₀ = z̄₀, ρ = ρ̄, ps = ps̄)

Reverse-mode derivative of $z^\star(z_0,\rho,\theta)$ by the adjoint method. Given the cotangent z̄ $= \partial\ell/\partial z^\star$ of a downstream loss, solve

\[J_{FF}^\top\lambda = -\bar z_F\]

once, and read every gradient off $\lambda$:

w.r.t.cotangent
parameters $\theta$$\sum_k w_k\lambda_{t_k}\,(\partial_\theta\varepsilon_\theta(x_k,t_k))^\top\lambda$ (via epsilon_vjp_params)
hard-clamped inputs $z_{0,C}$$\bar z_C + J_{FC}^\top\lambda$
soft targets $z_{0,F}$$-\rho_F^2\odot\lambda$
precisions $\rho_F$$2\rho_F\odot(z^\star_F - z_{0,F})\odot\lambda$

No noise has to be stored from the forward pass: the nodes are part of the model, and the Jacobians are evaluated at $z^\star$ afresh. Throws if sol did not converge, because the implicit function theorem says nothing about a point that is not a root (Backpropagation through Implicit Inference.md §6).

source
VariationalDiffusion.implicit_residual — Method
implicit_residual(m, z, z₀, ρ, ps, st) -> (r, st)

$r = g_\theta(z) + \rho^2\odot(z - z_0)$ on the free coordinates (finite ρ), zero on the hard-clamped ones ($\rho_i = \infty$), where the constraint is $z_i = z_{0,i}$ instead.

source
VariationalDiffusion.implicit_roots — Method
implicit_roots(m, z₀, ρ, ps, st; nstarts = 16, spread = 1.0, starts = nothing,
               rng = Xoshiro(0), unique_tol = 1e-3, kwargs...) -> (Vector{ImplicitSolution}, st)

All distinct stable answers to a query that a set of starting points reaches. Starts are z₀ itself and nstarts perturbations of it on the free coordinates (normal, standard deviation spread), or the vectors in starts; hard-clamped coordinates stay fixed. Each start is solved with implicit_infer (kwargs are passed on); converged, stable solutions are kept, and two are the same answer if they differ by less than unique_tol (relative).

The answers are ordered by the query's energy $U(z) + \tfrac12\lVert P(z - z_0)\rVert^2$ when the predictor is an EnergyNetwork (lowest first, i.e. the most plausible answer first), and otherwise in the order found. There is no guarantee that every stable answer is found; more starts and a larger spread find more.

source
VariationalDiffusion.implicit_solution — Method
implicit_solution(f::DiffusionFactor, polarity, inputs, π, ps, st) -> (ImplicitSolution, st)

Run the factor's implicit inversion and return the whole ImplicitSolution: the full state, the residual, and the converged / stable flags. invert and factor_message return only the target block as a DiracBelief; this is the call to make when the report matters.

The reference configuration is assemble_state(f, inputs, π): inputs on their channels, and the prior π's point (if any) on the others. It doubles as the warm start, which is how a message schedule reuses the previous belief: as in ImplicitLayers' DEQ factor, the prior seeds the solver. Whether it also pulls the solution is decided by the polarity's precision on the unobserved channels — default_precision(Unobserved()) == 1, an anchor; set it to 0 for a pure conditional.

source
VariationalDiffusion.integrated_beta — Method
integrated_beta(s::VPSDE, t) -> Real

$B(t) = \int_0^t \beta(u)\,du = \beta_{\min}t + \tfrac12(\beta_{\max}-\beta_{\min})t^2$.

In closed form, which is the only reason the VP-SDE is convenient: a schedule needing numerical quadrature here would put a solve inside every message.

source
VariationalDiffusion.marginal_variance — Method
marginal_variance(s, t, v₀) -> Real

The variance of $x_t$ when $x_0$ has variance v₀: $\alpha_t^2 v_0 + \sigma_t^2$.

For v₀ = 1 this is identically 1 — the variance-preserving property. It is used by the analytic Gaussian oracle in the test suite, which is the only closed-form $\varepsilon_\theta$ available and therefore the only thing the inversion can be checked against.

source
VariationalDiffusion.mixture_logdensity — Method
mixture_logdensity(l::GaussianMixtureEps, x, t, ps) -> Real

$\log p_t(x)$ of the mixture. Used to check that the deterministic RED-Diff field is the gradient of a smoothed log-density (Implicit Diffusion Learners.md §3).

source
VariationalDiffusion.noise_schedule — Method
noise_schedule(p::NoisePredictor) -> AbstractNoiseSchedule

The schedule the predictor was built with.

Named noise_schedule rather than schedule because the latter is Base.schedule (for Tasks), and shadowing it would be the sixth name collision in this project — see Mycelium.md §2 for the others.

source
VariationalDiffusion.noisefree_nodes — Method
noisefree_nodes(n, t)

One noise level, no noise: $g(z) = \lambda_t\,\varepsilon_\theta(\alpha_t z, t)$. This is the deterministic relaxation: a root satisfies $\hat x_0(\alpha_t z) = z$ on the free coordinates — a fixed point of the Tweedie denoiser, i.e. a deep equilibrium model whose layer is the denoiser (Deterministic Relaxation.md).

source
VariationalDiffusion.perturb — Method
perturb(s, x₀, t, ε) -> x_t

$x_t = \alpha_t x_0 + \sigma_t \varepsilon$: one draw from the perturbation kernel, written as a reparametrisation so it is differentiable in x₀.

That differentiability is exactly what RED-Diff keeps and what its stop-gradient does not discard — see reddiff.md §2.

source
VariationalDiffusion.precision_vector — Method
precision_vector(f, p::Polarity) -> (ρ, hardmask)

The diagonal of $P = \rho_{in}P_{in} + \rho_{out}P_{out} + \rho_{latent}P_{latent}$, as a vector of length n, together with a BitVector marking the coordinates whose precision is infinite.

Infinite entries are split off rather than stored as Inf: they cannot participate in a gradient, and the honest treatment is projection. default_precision(Observed()) == Inf, so by default every observed channel lands in hardmask and is clamped exactly.

source
VariationalDiffusion.prior_field — Method
prior_field(m, z, ps, st) -> (g, st)

$g_\theta(z) = \sum_k w_k\lambda_{t_k}\bigl(\varepsilon_\theta(\alpha_k z+\sigma_k\varepsilon_k,t_k)-\varepsilon_k\bigr)$.

source
VariationalDiffusion.prior_jacobian — Method
prior_jacobian(m, z, ps, st) -> Matrix

$\partial g/\partial z = \sum_k w_k\lambda_{t_k}\alpha_{t_k}\,\partial_x\varepsilon_\theta$. Symmetric iff the field is (locally) a gradient.

source
VariationalDiffusion.prox_infer — Method
prox_infer(p, z₀, ρ, t, ps, st; λ = 0.01, z_init = z₀, maxiters = 500, tol = 1e-9)
    -> (ImplicitSolution, st)

Deterministic inference for the implicit learner with a proximal prior instead of a score: half-quadratic splitting between the prior's prox at level t and the clamp's prox,

\[u \leftarrow \operatorname{prox}_{-\lambda\log p_t}(w),\qquad w_i \leftarrow \frac{u_i/\lambda + \rho_i^2 z_{0,i}}{1/\lambda + \rho_i^2}\quad(\rho_i<\infty),\qquad w_i \leftarrow z_{0,i}\quad(\rho_i=\infty).\]

A fixed point $(u, w)$ is a stationary point of $-\log p_t(u) + \tfrac1{2\lambda}\lVert u-w\rVert^2 + \tfrac12\lVert P(w - z_0)\rVert^2$ — the relaxed problem whose $\lambda\to0$ limit is the implicit learner's (Deterministic Relaxation.md); this is the plug-and-play / DiffPIR splitting. Hard clamps are exact. Returns w (which satisfies the clamps) as z, with residual $= \lVert u - w\rVert$; stable is not assessed by this solver and is reported true when converged.

source
VariationalDiffusion.proxdm_sample — Method
proxdm_sample(p, schedule, x, ps, st; steps = 50, hybrid = false, rng = Random.default_rng(),
              T = 1.0) -> (x₀, st)

ProxDM's sampler (Fang et al., Algorithm 1), started from x (pure noise for the full reverse process). With $\gamma_k = \int_{t_{k-1}}^{t_k}\beta(s)\,ds$ on a uniform grid $0 = t_0 < \dots < t_N = T$:

  • PDA (backward Euler): $X_{k-1} = \operatorname{prox}_{-\frac{2\gamma_k}{2-\gamma_k}\log p_{t_{k-1}}}\bigl(\tfrac{2}{2-\gamma_k}(X_k + \sqrt{\gamma_k}\,z_k)\bigr)$, which needs $\gamma_k < 2$;
  • PDA-hybrid (hybrid = true): $X_{k-1} = \operatorname{prox}_{-\gamma_k\log p_{t_{k-1}}}\bigl((1+\tfrac12\gamma_k)X_k + \sqrt{\gamma_k}\,z_k\bigr)$, no step-size limit.

Both follow from discretising the drift of the reverse VP-SDE at the new point (hybrid: the score term only), which turns each step into a proximal step. Unlike score-based samplers the last step denoises after adding noise, so no final denoising step is needed.

source
VariationalDiffusion.proximal — Function
proximal(p, v, t, λ, ps, st) -> (u, st)

$\operatorname{prox}_{-\lambda\log p_t}(v) = \arg\min_u \tfrac12\lVert u-v\rVert^2-\lambda\log p_t(u)$: a MAP denoiser at level t with regularisation λ. Its optimality condition is $u - v = \lambda\nabla\log p_t(u)$ — a backward (implicit) gradient step.

source
VariationalDiffusion.proximal_matching_loss — Method
proximal_matching_loss(p::ProxNetwork, x_t, t, λ, ε, ζ, ps, st) -> (ℓ, st)

One sample of ProxDM's training objective (their Eq. 9): $\ell_{PM}\bigl(\varepsilon_\theta(x_t + \sqrt\lambda\,\varepsilon; t, \lambda), \varepsilon; \zeta\bigr)$ with $\ell_{PM}(a, b; \zeta) = 1 - \exp\bigl(-\lVert a-b\rVert^2/(d\zeta^2)\bigr)$. As $\zeta\to0$ its minimiser is the prox (a MAP denoiser), where the squared loss would give the MMSE denoiser. Shrink $\zeta$ during training. x_t is a sample of $p_t$, ε a fresh standard normal.

source
VariationalDiffusion.reddiff_solve — Method
reddiff_solve(pred, cfg, x_init, datagrad, hard, ps, st) -> (x, st)

Algorithm 1: steps iterations of

\[x \leftarrow \mathrm{Optim}\bigl(x,\ \underbrace{\nabla_x \tfrac12\|P(x_0-x)\|^2}_{\texttt{datagrad}} + \underbrace{\mathbb{E}[\lambda_t(\varepsilon_\theta - \varepsilon)]}_{\text{Prop. 2}}\bigr)\]

datagrad(x) -> g supplies the clamp's gradient; keeping it a closure is what makes this operator reusable for a general forward model $f$ rather than only the diagonal $P$.

hard is a (mask, values) pair applied by projection after every step: coordinates with $\rho_{in} = \infty$ are overwritten rather than penalised, because an infinite precision is not a number you can put in a gradient. That projection is the categorical cup of Copiers Cups and Caps.md, and it is the one place the $\rho \to \infty$ limit is taken literally rather than numerically.

source
VariationalDiffusion.reddiff_weight — Method
reddiff_weight(cfg, s, t) -> Real

$\lambda_t = \lambda/\mathrm{SNR}_t = \lambda\,\sigma_t/\alpha_t$.

The paper motivates this as converting the noise-space objective to signal space. Note it grows without bound as $t \to 1$: late, heavily-noised times get the most weight, which is the opposite of the ELBO weighting and is why RED-Diff is mode-seeking rather than distribution-matching.

source
VariationalDiffusion.regulariser_gradient — Method
regulariser_gradient(pred, x, cfg, ps, st) -> (g, st)

$\nabla_x \mathrm{reg}(x) = \mathbb{E}_{t,\varepsilon}[\lambda_t(\varepsilon_\theta(x_t,t) - \varepsilon)]$ — Proposition 2, by Monte Carlo with cfg.samples draws.

This is the entire prior contribution: one forward pass of the network per draw, no Jacobian, no backward pass. The estimator is unbiased in $(t,\varepsilon)$ but the underlying gradient is not the true gradient of the regulariser — the denoiser Jacobian has been dropped. reddiff.md §2 is about what that costs.

source
VariationalDiffusion.score — Method
score(p::NoisePredictor, x, t, ps, st) -> (s, st)

$s_\theta(x,t) = -\varepsilon_\theta(x,t)/\sigma_t \approx \nabla_x \log p_t(x)$.

The identity is exact for the perturbation kernel: with $x_t = \alpha_t x_0 + \sigma_t\varepsilon$,

\[\nabla_{x_t}\log p_{0t}(x_t\mid x_0) = -\frac{x_t - \alpha_t x_0}{\sigma_t^2} = -\frac{\varepsilon}{\sigma_t}\]

so a perfect noise predictor is a perfect score. The $1/\sigma_t$ is why VPSDE carries a tmin.

source
VariationalDiffusion.sigma — Method
sigma(s, t) -> Real

The noise coefficient $\sigma_t = \sqrt{1 - \exp(-B(t))} = \sqrt{1 - \alpha_t^2}$.

Computed from expm1 rather than as sqrt(1 - alpha^2): near t = 0 the latter is sqrt(1 - (1-ε)²) and loses half its significant digits to cancellation.

source