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.jlTheory: 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.
| type | what it is |
|---|---|
AbstractProximalPredictor | anything 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_samplethrows 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):
| steps | 20 | 40 | 100 | 400 |
|---|---|---|---|---|
| sample std (true 0.30) | 0.248 | 0.268 | 0.286 | 0.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.
MixtureProxuses Newton where is positive definite and a gradient step otherwise, and returns the local minimiser reached from . A learnedProxNetworkmakes 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_inferyet. The fixed point satisfies , so the implicit function theorem applies just as forimplicit_infer. It needs the prox’s input Jacobian and parameter VJP, which is whyProxNetworkwould get anadfield when this is added (backends). stableis not assessed by the splitting; it reportstruewhen converged.- A
ProxDMfactor inversion (alongsideREDDiffandImplicitProxinfactor.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