implementation

Proximal diffusion models: a diffusion prior that is queried through its proximal operator instead of its score. Implemented here are the interface, an exact oracle, ProxDM’s sampler, proximal inference for the implicit learner, and the training loss.

Sources: Fang, Díaz, Buchanan & Sulam, Beyond Scores: Proximal Diffusion Models, arXiv:2507.08956 (Algorithm 1, Eq. 9); Zhu et al., Denoising Diffusion Models for Plug-and-Play Image Restoration, CVPR 2023 (DiffPIR, the half-quadratic splitting); code: proxdm.jl

Theory: ProxDM and Proximal Alternatives · Deterministic Relaxation · Inference Signatures

Notation. As everywhere in this vault, is the joint state and the prior is on . The paper writes for the same thing.

1. The object

The optimality condition is a gradient step evaluated at the new point: a backward (implicit) Euler step, where a score is a forward one. That is the whole idea of ProxDM: discretise the reverse SDE backwards, and the network you need is a prox, not a score.

typewhat it is
AbstractProximalPredictoranything with proximal(p, v, t, λ, ps, st) -> (u, st)
ProxNetwork(model, schedule; input)a Lux model in ProxDM’s parametrisation ; conditioned on two scalars
MixtureProx(::GaussianMixtureEps)the exact prox of a Gaussian mixture, by Newton with backtracking; the oracle a perfectly trained ProxNetwork would match

MixtureProx uses and its Hessian directly from the mixture, not , so it is valid at , where the sampler’s last step needs it.

2. The sampler (Algorithm 1)

On a uniform grid with :

  • PDA (hybrid = false): , which needs ; proxdm_sample throws otherwise.
  • PDA-hybrid (hybrid = true): : the linear drift explicit, the score implicit; no step-size limit.

The last step evaluates the prox at , so the output is already denoised; there is no separate final denoising step.

Validation (one Gaussian, mean , std , exact prox, 1,500 samples):

steps2040100400
sample std (true 0.30)0.2480.2680.2860.293

The mean is right at every step count; the spread converges at first order, as a backward Euler scheme should. The sampler on the circle mixture gives radius .

3. Proximal inference for the implicit learner

prox_infer(p, z₀, ρ, t, ps, st; λ) is the implicit learner’s query with a proximal prior. It splits the relaxed problem

and alternates the two proximal maps (half-quadratic splitting, as in DiffPIR):

Hard clamps are exact; the free coordinates () take . As the coupling becomes and the problem becomes the implicit learner’s (Deterministic Relaxation). It returns an ImplicitSolution with z = w (which satisfies the clamps) and residual .

On the circle mixture, the query ”, which ?” from returns , symmetric to the last digit, in 20 iterations.

4. Training: proximal matching

proximal_matching_loss is one sample of the paper’s Eq. 9:

The squared loss would train the MMSE denoiser ; this bounded loss, as , trains the MAP denoiser, which is the prox. Shrink over training. The loss is a plain function of ps, so any Lux training loop and any AD backend trains it.

5. Implementation difficulties and open ends

  • Non-convex priors make the prox multivalued. For a mixture, is not convex, and can have several minimisers. MixtureProx uses Newton where is positive definite and a gradient step otherwise, and returns the local minimiser reached from . A learned ProxNetwork makes the same choice implicitly, wherever training put it.
  • The PDA step-size limit. fails for coarse grids near under the default VP schedule. The hybrid scheme removes it; PDA throws with a message saying so instead of returning garbage.
  • No adjoint for prox_infer yet. The fixed point satisfies , so the implicit function theorem applies just as for implicit_infer. It needs the prox’s input Jacobian and parameter VJP, which is why ProxNetwork would get an ad field when this is added (backends).
  • stable is not assessed by the splitting; it reports true when converged.
  • A ProxDM factor inversion (alongside REDDiff and ImplicitProx in factor.jl) is the natural next step and would make ProxDM usable inside a factor graph. It is the first case with two parameter trees (the prior’s and the prox network), the structure ProxDM and Proximal Alternatives §4 anticipated.

Related: implicit, reddiff, backends, analytic, ProxDM and Proximal Alternatives, Inference Signatures, Deterministic Relaxation