Slide 1

Slide 1 text

Posterior Sampling with the nonconvex Proximal Stochastic Gradient Langevin Algorithm (PSGLA) Marien Renaud, Valentin de Bortoli, Arthur Leclaire, Nicolas Papadakis

Slide 2

Slide 2 text

Outline 1 Sampling with Langevin algorithms 2 Stability result for PSGLA 3 Numerical illustrations

Slide 3

Slide 3 text

1. Sampling with Langevin algorithms

Slide 4

Slide 4 text

Inverse problems in imaging Acquisition Inverse problem x ∈ Rn y = Noise(Ax) ∈ Rm 1/18

Slide 5

Slide 5 text

Inverse problems in imaging Problem: • Recover x ∗ ∈ Rn from a degraded observation y = Ax ∗ + η, for a linear operator A: Rn → Rd and η ∼ N (0, σ 2 Id ) 2/18

Slide 6

Slide 6 text

Inverse problems in imaging Problem: • Recover x ∗ ∈ Rn from a degraded observation y = Ax ∗ + η, for a linear operator A: Rn → Rd and η ∼ N (0, σ 2 Id ) 1 2 • Data fidelity term p(y |x) ∝ e − 2 ||Ax−y || = e −f (x) 2/18

Slide 7

Slide 7 text

Inverse problems in imaging Problem: • Recover x ∗ ∈ Rn from a degraded observation y = Ax ∗ + η, for a linear operator A: Rn → Rd and η ∼ N (0, σ 2 Id ) 1 2 • Data fidelity term p(y |x) ∝ e − 2 ||Ax−y || = e −f (x) p(x) • Regularization required: prior on image manifold p(x) ∝ e −g(x) 2/18

Slide 8

Slide 8 text

Inverse problems in imaging Problem: • Recover x ∗ ∈ Rn from a degraded observation y = Ax ∗ + η, for a linear operator A: Rn → Rd and η ∼ N (0, σ 2 Id ) 1 2 • Data fidelity term p(y |x) ∝ e − 2 ||Ax−y || = e −f (x) p(x) • Regularization required: prior on image manifold p(x) ∝ e −g(x) Resolution x̂ ∈ arg max p(x|y ) ∝ p(y |x)p(x) x 2/18

Slide 9

Slide 9 text

Inverse problems in imaging Problem: • Recover x ∗ ∈ Rn from a degraded observation y = Ax ∗ + η, for a linear operator A: Rn → Rd and η ∼ N (0, σ 2 Id ) 1 2 • Data fidelity term p(y |x) ∝ e − 2 ||Ax−y || = e −f (x) p(x) • Regularization required: prior on image manifold p(x) ∝ e −g(x) Resolution x̂ ∈ arg max p(x|y ) ∝ p(y |x)p(x) x ∈ arg min − log p(y |x) − log p(x) | {z } | {z } x f (x) g(x) 2/18

Slide 10

Slide 10 text

Inverse problems in imaging Problem: • Recover x ∗ ∈ Rn from a degraded observation y = Ax ∗ + η, for a linear operator A: Rn → Rd and η ∼ N (0, σ 2 Id ) 1 2 • Data fidelity term p(y |x) ∝ e − 2 ||Ax−y || = e −f (x) p(x) • Regularization required: prior on image manifold p(x) ∝ e −g(x) Resolution x̂ ∈ arg max p(x|y ) ∝ p(y |x)p(x) x ∈ arg min − log p(y |x) − log p(x) | {z } | {z } x f (x) g(x) ∈ arg min V (x) := f (x) + g(x) x 2/18

Slide 11

Slide 11 text

Inverse problems in imaging Problem: • Recover x ∗ ∈ Rn from a degraded observation y = Ax ∗ + η, for a linear operator A: Rn → Rd and η ∼ N (0, σ 2 Id ) 1 2 • Data fidelity term p(y |x) ∝ e − 2 ||Ax−y || = e −f (x) p(x) • Regularization required: prior on image manifold p(x) ∝ e −g(x) Resolution x̂ ∈ arg max p(x|y ) ∝ p(y |x)p(x) x ∈ arg min − log p(y |x) − log p(x) | {z } | {z } x f (x) g(x) ∈ arg min V (x) := f (x) + g(x) x ∈ arg max e −V (x) x 2/18

Slide 12

Slide 12 text

Inverse problems and Posterior sampling 3/18

Slide 13

Slide 13 text

Inverse problems and Posterior sampling 3/18

Slide 14

Slide 14 text

Inverse problems and Posterior sampling 3/18

Slide 15

Slide 15 text

Inverse problems and Posterior sampling 3/18

Slide 16

Slide 16 text

Inverse problems and Posterior sampling 3/18

Slide 17

Slide 17 text

Inverse problems and Posterior sampling 3/18

Slide 18

Slide 18 text

Inverse problems and Posterior sampling 3/18

Slide 19

Slide 19 text

Inverse problems and Posterior sampling 3/18

Slide 20

Slide 20 text

Inverse problems and Posterior sampling 3/18

Slide 21

Slide 21 text

Inverse problems and Posterior sampling 3/18

Slide 22

Slide 22 text

Inverse problems and Posterior sampling 3/18

Slide 23

Slide 23 text

Inverse problems and Posterior sampling p(x|y ) ∝ p(y |x)p(x) ∝ e −f (x)−g(x) ∝ e −V (x) 3/18

Slide 24

Slide 24 text

Posterior sampling with Langevin Algorithms Problem: How to sample a probability distribution π ∝ e −V (π not log concave ⇔ V non-convex)? 4/18

Slide 25

Slide 25 text

Posterior sampling with Langevin Algorithms Problem: How to sample a probability distribution π ∝ e −V (π not log concave ⇔ V non-convex)? In Imaging: • Gibbs sampling [Vono et al. ’22, Coeurdoux et al. ’24, Kuric et al. ’25, Bouton et al. ’26] • Langevin [Pereyra’16, Durmus et al.’18, Luu et al.’21, Laumont et al.’22, Klatzer et al.’25, Duan et al.’26] • Diffusion [Song et al. ’21, Yismaw et al. ’25] 4/18

Slide 26

Slide 26 text

Posterior sampling with Langevin Algorithms Problem: How to sample a probability distribution π ∝ e −V (π not log concave ⇔ V non-convex)? In Imaging: • Gibbs sampling [Vono et al. ’22, Coeurdoux et al. ’24, Kuric et al. ’25, Bouton et al. ’26] • Langevin [Pereyra’16, Durmus et al.’18, Luu et al.’21, Laumont et al.’22, Klatzer et al.’25, Duan et al.’26] • Diffusion [Song et al. ’21, Yismaw et al. ’25] 4/18

Slide 27

Slide 27 text

Posterior sampling with Langevin Algorithms Problem: How to sample a probability distribution π ∝ e −V (π not log concave ⇔ V non-convex)? In Imaging: • Gibbs sampling [Vono et al. ’22, Coeurdoux et al. ’24, Kuric et al. ’25, Bouton et al. ’26] • Langevin [Pereyra’16, Durmus et al.’18, Luu et al.’21, Laumont et al.’22, Klatzer et al.’25, Duan et al.’26] • Diffusion [Song et al. ’21, Yismaw et al. ’25] Langevin stochastic differential equation [Roberts and Tweedie ’96], with a Brownian motion wt √ dxt = −∇V (xt )dt + 2dwt 4/18

Slide 28

Slide 28 text

Posterior sampling with Langevin Algorithms Problem: How to sample a probability distribution π ∝ e −V (π not log concave ⇔ V non-convex)? In Imaging: • Gibbs sampling [Vono et al. ’22, Coeurdoux et al. ’24, Kuric et al. ’25, Bouton et al. ’26] • Langevin [Pereyra’16, Durmus et al.’18, Luu et al.’21, Laumont et al.’22, Klatzer et al.’25, Duan et al.’26] • Diffusion [Song et al. ’21, Yismaw et al. ’25] Langevin stochastic differential equation [Roberts and Tweedie ’96], with a Brownian motion wt √ dxt = −∇V (xt )dt + 2dwt ✓ xt admits π as invariant distribution: pxt ∝ π 4/18

Slide 29

Slide 29 text

Posterior sampling with Langevin Algorithms Problem: How to sample a probability distribution π ∝ e −V (π not log concave ⇔ V non-convex)? In Imaging: • Gibbs sampling [Vono et al. ’22, Coeurdoux et al. ’24, Kuric et al. ’25, Bouton et al. ’26] • Langevin [Pereyra’16, Durmus et al.’18, Luu et al.’21, Laumont et al.’22, Klatzer et al.’25, Duan et al.’26] • Diffusion [Song et al. ’21, Yismaw et al. ’25] Langevin stochastic differential equation [Roberts and Tweedie ’96], with a Brownian motion wt √ dxt = −∇V (xt )dt + 2dwt ✓ xt admits π as invariant distribution: pxt ∝ π ✓ Uniform measure supported on the iterates (Xk )0≤k≤N used as an estimator of π: pXk ∝ ∼π 4/18

Slide 30

Slide 30 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient 5/18

Slide 31

Slide 31 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Gradient Descent: xk+1 = xk − γ∇V (xk ) 5/18

Slide 32

Slide 32 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Gradient Descent: xk+1 = xk − γ∇V (xk ) 5/18

Slide 33

Slide 33 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Gradient Descent: xk+1 = xk − γ∇V (xk ) 5/18

Slide 34

Slide 34 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Gradient Descent: xk+1 = xk − γ∇V (xk ) 5/18

Slide 35

Slide 35 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Gradient Descent: xk+1 = xk − γ∇V (xk ) 5/18

Slide 36

Slide 36 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Gradient Descent: xk+1 = xk − γ∇V (xk ) 5/18

Slide 37

Slide 37 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Gradient Descent: xk+1 = xk − γ∇V (xk ) 5/18

Slide 38

Slide 38 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 39

Slide 39 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 40

Slide 40 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 41

Slide 41 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 42

Slide 42 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 43

Slide 43 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 44

Slide 44 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 45

Slide 45 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 46

Slide 46 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Stochastic Gradient Descent: xk+1 = xk − γ∇V (xk ) + γzk+1 5/18

Slide 47

Slide 47 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 48

Slide 48 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 49

Slide 49 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 50

Slide 50 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 51

Slide 51 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 52

Slide 52 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 53

Slide 53 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 54

Slide 54 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 55

Slide 55 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 56

Slide 56 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 57

Slide 57 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 58

Slide 58 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 59

Slide 59 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 60

Slide 60 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 61

Slide 61 text

Optimizing vs Sampling Assume V non-convex with Lipschitz gradient Langevin: √ xk+1 = xk − γ∇V (xk ) + 2γzk+1 5/18

Slide 62

Slide 62 text

Posterior sampling with Langevin Algorithms Langevin stochastic differential equation [Roberts and Tweedie ’96], with a Brownian motion wt √ dxt = −∇V (xt )dt + 2dwt 6/18

Slide 63

Slide 63 text

Posterior sampling with Langevin Algorithms Langevin stochastic differential equation [Roberts and Tweedie ’96], with a Brownian motion wt √ dxt = −∇V (xt )dt + 2dwt ✗ ∇V = −∇ log π is often unknown 6/18

Slide 64

Slide 64 text

Posterior sampling with Langevin Algorithms Langevin stochastic differential equation [Roberts and Tweedie ’96], with a Brownian motion wt √ dxt = −∇V (xt )dt + 2dwt ✗ ∇V = −∇ log π is often unknown Implementation: inexact Unadjusted Langevin Algorithm (iULA) • Use of a drift function b ≈ ∇V (inexact) • Euler-Maruyama discretization, for X0 ∈ Rn and Zk+1 ∼ N (0, Id ) p Xk+1 = Xk − γb(Xk ) + 2γZk+1 , with γ > 0 (unadjusted) 6/18

Slide 65

Slide 65 text

Posterior sampling with Langevin Algorithms Langevin stochastic differential equation [Roberts and Tweedie ’96], with a Brownian motion wt √ dxt = −∇V (xt )dt + 2dwt ✗ ∇V = −∇ log π is often unknown Implementation: inexact Unadjusted Langevin Algorithm (iULA) • Use of a drift function b ≈ ∇V (inexact) • Euler-Maruyama discretization, for X0 ∈ Rn and Zk+1 ∼ N (0, Id ) p Xk+1 = Xk − γb(Xk ) + 2γZk+1 , with γ > 0 (unadjusted) Why can we use the uniform measure supported on (Xk )0≤k≤N as an estimator of π? 6/18

Slide 66

Slide 66 text

Posterior sampling with Langevin Algorithms Convergence of inexact Unadjusted Langevin Algorithm (iULA) p Xk+1 = Xk − γb(Xk ) + 2γZk+1 7/18

Slide 67

Slide 67 text

Posterior sampling with Langevin Algorithms Convergence of inexact Unadjusted Langevin Algorithm (iULA) p Xk+1 = Xk − γb(Xk ) + 2γZk+1 Two important notions 1 Invariant law: there exists µ such that if X0 ∼ µ, then ∀k ≥ 0, Xk ∼ µ 7/18

Slide 68

Slide 68 text

Posterior sampling with Langevin Algorithms Convergence of inexact Unadjusted Langevin Algorithm (iULA) p Xk+1 = Xk − γb(Xk ) + 2γZk+1 Two important notions 1 Invariant law: there exists µ such that if X0 ∼ µ, then ∀k ≥ 0, Xk ∼ µ 2 Geometric ergodicity: For pXk the law of Xk , ∃A ≥ 0, ρ ∈ (0, 1) and an invariant law µ W1 (pXk , µ) ≤ Aρk 7/18

Slide 69

Slide 69 text

Posterior sampling with Langevin Algorithms Convergence of inexact Unadjusted Langevin Algorithm (iULA) p Xk+1 = Xk − γb(Xk ) + 2γZk+1 Two important notions 1 Invariant law: there exists µ such that if X0 ∼ µ, then ∀k ≥ 0, Xk ∼ µ 2 Geometric ergodicity: For pXk the law of Xk , ∃A ≥ 0, ρ ∈ (0, 1) and an invariant law µ W1 (pXk , µ) ≤ Aρk Sufficient conditions on the drift b for Xk to be geometrically ergodic: • b is L-Lipschitz, i.e. ∀x, y ∈ Rd , ∥b(x) − b(y )∥ ≤ L∥x − y ∥ • ∃R, m > 0 such that ∀x, y ∈ Rd with ∥x − y ∥ ≥ R, ⟨b(x) − b(y ), x − y ⟩ ≥ m∥x − y ∥2 7/18

Slide 70

Slide 70 text

Inverse problems and composite potentials • Inverse problems: V = f + g Xk+1 = Xk − γ(∇f (Xk ) + ∇g(Xk )) + p 2γZk+1 ✓ f (x) = 12 ||Ax − y ||2 ✗ g(x) = − log p(x) is unknown 8/18

Slide 71

Slide 71 text

Inverse problems and composite potentials • Inverse problems: V = f + g Xk+1 = Xk − γ(∇f (Xk ) + ∇g(Xk )) + p 2γZk+1 ✓ f (x) = 12 ||Ax − y ||2 ✗ g(x) = − log p(x) is unknown • RED [Romano et al. ’17]: smoothed prior g = − log pϵ , pϵ = p ⋆ N (0, ϵ2 Id ) and score ∇g estimated with a MMSE denoiser Dϵ Xk − Dϵ (Xk ) ∇g(Xk ) = −∇ log pϵ (Xk ) = (Tweedie formula) ϵ2 8/18

Slide 72

Slide 72 text

Inverse problems and composite potentials • Inverse problems: V = f + g Xk+1 = Xk − γ(∇f (Xk ) + ∇g(Xk )) + p 2γZk+1 ✓ f (x) = 12 ||Ax − y ||2 ✗ g(x) = − log p(x) is unknown • RED [Romano et al. ’17]: smoothed prior g = − log pϵ , pϵ = p ⋆ N (0, ϵ2 Id ) and score ∇g estimated with a MMSE denoiser Dϵ Xk − Dϵ (Xk ) ∇g(Xk ) = −∇ log pϵ (Xk ) = (Tweedie formula) ϵ2 • PnP-ULA [Laumont et al. ’22]   p 1 Xk+1 = Xk − γ ∇f (Xk ) + 2 (Xk − Dϵ (Xk )) + 2γZk+1 ϵ 8/18

Slide 73

Slide 73 text

Inverse problems and composite potentials • Inverse problems: V = f + g Xk+1 = Xk − γ(∇f (Xk ) + ∇g(Xk )) + p 2γZk+1 ✓ f (x) = 12 ||Ax − y ||2 ✗ g(x) = − log p(x) is unknown • RED [Romano et al. ’17]: smoothed prior g = − log pϵ , pϵ = p ⋆ N (0, ϵ2 Id ) and score ∇g estimated with a MMSE denoiser Dϵ Xk − Dϵ (Xk ) ∇g(Xk ) = −∇ log pϵ (Xk ) = (Tweedie formula) ϵ2 • PnP-ULA [Laumont et al. ’22]   p 1 Xk+1 = Xk − γ ∇f (Xk ) + 2 (Xk − Dϵ (Xk )) + 2γZk+1 ϵ ✓ Sampling stability results for non-convex g 8/18

Slide 74

Slide 74 text

Inverse problems and composite potentials • Inverse problems: V = f + g Xk+1 = Xk − γ(∇f (Xk ) + ∇g(Xk )) + p 2γZk+1 ✓ f (x) = 12 ||Ax − y ||2 ✗ g(x) = − log p(x) is unknown • RED [Romano et al. ’17]: smoothed prior g = − log pϵ , pϵ = p ⋆ N (0, ϵ2 Id ) and score ∇g estimated with a MMSE denoiser Dϵ Xk − Dϵ (Xk ) ∇g(Xk ) = −∇ log pϵ (Xk ) = (Tweedie formula) ϵ2 • PnP-ULA [Laumont et al. ’22]   p 1 Xk+1 = Xk − γ ∇f (Xk ) + 2 (Xk − Dϵ (Xk )) + 2γZk+1 ϵ ✓ Sampling stability results for non-convex g ✗ Samples Xk may be noisy, long mixing time 8/18

Slide 75

Slide 75 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + p 2γZk+1 9/18

Slide 76

Slide 76 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + p (Id + γ∇g) (Xk+1 ) = Xk − γ∇f (Xk ) + 2γZk+1 p 2γZk+1 9/18

Slide 77

Slide 77 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: p Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + 2γZk+1 p (Id + γ∇g) (Xk+1 ) = Xk − γ∇f (Xk ) + 2γZk+1   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 (PSGLA) 9/18

Slide 78

Slide 78 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: p Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + 2γZk+1 p (Id + γ∇g) (Xk+1 ) = Xk − γ∇f (Xk ) + 2γZk+1   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 (PSGLA) - Studied for g convex [Salim and Richtarik ’20, Ehrhardt et al. ’24] 9/18

Slide 79

Slide 79 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: p Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + 2γZk+1 p (Id + γ∇g) (Xk+1 ) = Xk − γ∇f (Xk ) + 2γZk+1   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 (PSGLA) - Studied for g convex [Salim and Richtarik ’20, Ehrhardt et al. ’24] - Related to MYULA [Durmus et al. ’22], DAZ [Habring et al. ’25] 9/18

Slide 80

Slide 80 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: p Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + 2γZk+1 p (Id + γ∇g) (Xk+1 ) = Xk − γ∇f (Xk ) + 2γZk+1   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 (PSGLA) - Studied for g convex [Salim and Richtarik ’20, Ehrhardt et al. ’24] - Related to MYULA [Durmus et al. ’22], DAZ [Habring et al. ’25] • Plug-and-Play [Venkatakrishnan et al. ‘13]: Use a MAP denoiser Dγ ≈ Proxγg = Prox−γ log p 9/18

Slide 81

Slide 81 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: p Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + 2γZk+1 p (Id + γ∇g) (Xk+1 ) = Xk − γ∇f (Xk ) + 2γZk+1   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 (PSGLA) - Studied for g convex [Salim and Richtarik ’20, Ehrhardt et al. ’24] - Related to MYULA [Durmus et al. ’22], DAZ [Habring et al. ’25] • Plug-and-Play [Venkatakrishnan et al. ‘13]: Use a MAP denoiser Dγ ≈ Proxγg = Prox−γ log p ✗ State-of-the-art Prox denoisers correspond to weakly convex1 potentials g [Hurault et al. ’22] 1 ∃ρ > 0 such that g(.) + 2ρ ||.||2 is convex 9/18

Slide 82

Slide 82 text

Proximal Stochastic Gradient Langevin Algorithm (PSGLA) • Semi implicit discretization: p Xk+1 = Xk − γ∇f (Xk ) − γ∇g(Xk+1 ) + 2γZk+1 p (Id + γ∇g) (Xk+1 ) = Xk − γ∇f (Xk ) + 2γZk+1   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 (PSGLA) - Studied for g convex [Salim and Richtarik ’20, Ehrhardt et al. ’24] - Related to MYULA [Durmus et al. ’22], DAZ [Habring et al. ’25] • Plug-and-Play [Venkatakrishnan et al. ‘13]: Use a MAP denoiser Dγ ≈ Proxγg = Prox−γ log p ✗ State-of-the-art Prox denoisers correspond to weakly convex1 potentials g [Hurault et al. ’22] Our work: study the stability of PSGLA for non-convex potentials g 1 ∃ρ > 0 such that g(.) + 2ρ ||.||2 is convex 9/18

Slide 83

Slide 83 text

2. Stability of PSGLA

Slide 84

Slide 84 text

Sampling stability: Problem statement Sampling π ∝ e −f −g using PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 with pXk the distribution of Xk 10/18

Slide 85

Slide 85 text

Sampling stability: Problem statement Sampling π ∝ e −f −g using PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 with pXk the distribution of Xk Stability: Check if pXk ∝ ∼ π by studying the p-Wasserstein distance Wp (pXk , π)   p1 Z ∥x − y ∥p dβ(x, y ) Wp (µ, ν) = min β∈Π d Rd ×Rd d with Π the set of probability law β on R × R with marginals µ and ν 10/18

Slide 86

Slide 86 text

Sampling stability: Problem statement Sampling π ∝ e −f −g using PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 with pXk the distribution of Xk Stability: Check if pXk ∝ ∼ π by studying the p-Wasserstein distance Wp (pXk , π)   p1 Z ∥x − y ∥p dβ(x, y ) Wp (µ, ν) = min β∈Π d Rd ×Rd d with Π the set of probability law β on R × R with marginals µ and ν (Lazy) strategy: Build on top of PnP-ULA results [Laumont et al. ’22] p X̃k+1 = X̃k − γ(∇f (X̃k ) + g(X̃k )) + 2γZk+1 for which W1 (pX̃k , π) is controlled 10/18

Slide 87

Slide 87 text

Reformulation of PSGLA PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 1. Shadow sequence Yk+1 = Xk − γ∇f (Xk ) + p 2γZk+1 Xk+1 = Proxγg (Yk+1 ) thus Yk+1 = Proxγg (Yk ) − γ∇f (Proxγg (Yk )) + p 2γZk+1 11/18

Slide 88

Slide 88 text

Reformulation of PSGLA PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 1. Shadow sequence p Yk+1 = Xk − γ∇f (Xk ) + 2γZk+1 2. Moreau envelope g γ (y ) = infn z∈R Xk+1 = Proxγg (Yk+1 ) thus p Yk+1 = Proxγg (Yk ) − γ∇f (Proxγg (Yk )) + 2γZk+1 1 ∥y − z∥2 + g(z) 2γ If g is ρ-weakly convex and ργ < 1 then Proxγg (y ) = y − γ∇g γ (y ) 11/18

Slide 89

Slide 89 text

Reformulation of PSGLA PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 1. Shadow sequence 2. Moreau envelope p Yk+1 = Xk − γ∇f (Xk ) + 2γZk+1 g γ (y ) = infn z∈R Xk+1 = Proxγg (Yk+1 ) 1 ∥y − z∥2 + g(z) 2γ If g is ρ-weakly convex and ργ < 1 then thus Proxγg (y ) = y − γ∇g γ (y ) p Yk+1 = Proxγg (Yk ) − γ∇f (Proxγg (Yk )) + 2γZk+1 • The shadow sequence satisfies Yk+1 = Yk − γbγ (Yk ) + γ γ p 2γZk+1 , γ with the drift b (y ) = ∇f (y −γ∇g (y )) + ∇g (y ) 11/18

Slide 90

Slide 90 text

Reformulation of PSGLA PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 1. Shadow sequence 2. Moreau envelope p Yk+1 = Xk − γ∇f (Xk ) + 2γZk+1 g γ (y ) = infn z∈R Xk+1 = Proxγg (Yk+1 ) 1 ∥y − z∥2 + g(z) 2γ If g is ρ-weakly convex and ργ < 1 then thus Proxγg (y ) = y − γ∇g γ (y ) p Yk+1 = Proxγg (Yk ) − γ∇f (Proxγg (Yk )) + 2γZk+1 • The shadow sequence satisfies Yk+1 = Yk − γbγ (Yk ) + γ γ p 2γZk+1 , γ with the drift b (y ) = ∇f (y −γ∇g (y )) + ∇g (y ) 11/18

Slide 91

Slide 91 text

Reformulation of PSGLA PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 1. Shadow sequence 2. Moreau envelope p Yk+1 = Xk − γ∇f (Xk ) + 2γZk+1 g γ (y ) = infn z∈R Xk+1 = Proxγg (Yk+1 ) 1 ∥y − z∥2 + g(z) 2γ If g is ρ-weakly convex and ργ < 1 then thus Proxγg (y ) = y − γ∇g γ (y ) p Yk+1 = Proxγg (Yk ) − γ∇f (Proxγg (Yk )) + 2γZk+1 • The shadow sequence satisfies Yk+1 = Yk − γbγ (Yk ) + γ γ p 2γZk+1 , γ with the drift b (y ) = ∇f (y −γ∇g (y )) + ∇g (y ) • MYULA [Durmus et al. ’22] and DAZ [Habring et al. ’25]: b̃γ (y ) = ∇f (y ) + ∇g γ (y ) 11/18

Slide 92

Slide 92 text

Reformulation of PSGLA PSGLA   p Xk+1 = Proxγg Xk − γ∇f (Xk ) + 2γZk+1 1. Shadow sequence 2. Moreau envelope p Yk+1 = Xk − γ∇f (Xk ) + 2γZk+1 g γ (y ) = infn z∈R Xk+1 = Proxγg (Yk+1 ) 1 ∥y − z∥2 + g(z) 2γ If g is ρ-weakly convex and ργ < 1 then thus Proxγg (y ) = y − γ∇g γ (y ) p Yk+1 = Proxγg (Yk ) − γ∇f (Proxγg (Yk )) + 2γZk+1 • The shadow sequence satisfies Yk+1 = Yk − γbγ (Yk ) + γ γ p 2γZk+1 , γ with the drift b (y ) = ∇f (y −γ∇g (y )) + ∇g (y ) • MYULA [Durmus et al. ’22] and DAZ [Habring et al. ’25]: b̃γ (y ) = ∇f (y ) + ∇g γ (y ) • Postulate: the distribution of Yk should be close to that of µγ ∝ e −f −g γ 11/18

Slide 93

Slide 93 text

Stability of PSGLA PSGLA Yk+1 = Yk − γ (∇f (y −γ∇g γ (y )) + ∇g γ (y )) + p 2γZk+1 Xk+1 = Proxγg (Yk+1 ) Assumptions For 0 < γ < 1/ρ • ∇f is Lf Lipschitz • g γ is strongly convex at infinity (∇2 g γ ⪰ µId ) • g is ρ-weakly convex • g is Lg smooth on Proxγg (Rn ) Theorem There exist C1 , C2 , C3 , C4 ∈ R+ , r ∈ (0, 1) and γ̄ such that ∀γ ≤ γ̄ 1 Wp (pYk , µγ ) ≤ C1 r kγ + C2 γ 2p γ with pYk the distribution of Yk and µγ ∝ e −f −g , 1 Wp (pXk , νγ ) ≤ C3 r kγ + C4 γ 2p γ with pXk the distribution of Xk and νγ ∝ Proxγg #e −f −g . 12/18

Slide 94

Slide 94 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 with invariant laws πγi , we have 1 Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓp (π1 ) 2 γ 13/18

Slide 95

Slide 95 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: 1 p Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 Xk+1 = Xk − γ∇V (Xk ) + p 2γZk+1 1 Wp (πγ , π) ≤ C γ 2p γ 13/18

Slide 96

Slide 96 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: 1 p Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 Xk+1 = Xk − γ∇V (Xk ) + p 2γZk+1 1 Wp (πγ , π) ≤ C γ 2p γ Generalization of [Brosse et al. ’19, Renaud et al. ’24] to weakly convex functions 13/18

Slide 97

Slide 97 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: 1 p Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 Xk+1 = Xk − γ∇V (Xk ) + p 2γZk+1 1 Wp (πγ , π) ≤ C γ 2p γ Generalization of [Brosse et al. ’19, Renaud et al. ’24] to weakly convex functions √ • Chain Yk+1 = Yk − γbγ (Yk ) + 2γZk+1 , with bγ (y ) = ∇f (y − γ∇g γ (y )) + ∇g γ (y ) 13/18

Slide 98

Slide 98 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: 1 p Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 Xk+1 = Xk − γ∇V (Xk ) + p 2γZk+1 1 Wp (πγ , π) ≤ C γ 2p γ Generalization of [Brosse et al. ’19, Renaud et al. ’24] to weakly convex functions √ • Chain Yk+1 = Yk − γbγ (Yk ) + 2γZk+1 , with bγ (y ) = ∇f (y − γ∇g γ (y )) + ∇g γ (y ) - Geometric ergodicity: Wp (pYk , p∞ ) ≤ Ar kγ 13/18

Slide 99

Slide 99 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: Xk+1 = Xk − γ∇V (Xk ) + 1 p 1 Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 p 2γZk+1 Wp (πγ , π) ≤ C γ 2p γ Generalization of [Brosse et al. ’19, Renaud et al. ’24] to weakly convex functions √ • Chain Yk+1 = Yk − γbγ (Yk ) + 2γZk+1 , with bγ (y ) = ∇f (y − γ∇g γ (y )) + ∇g γ (y ) - Geometric ergodicity: Wp (pYk , p∞ ) ≤ Ar kγ 1 - Drift: Wp (p∞ , pγ ) ≤ B∥bγ − ∇(f + g γ )∥ℓp (π1 ) 2 γ 13/18

Slide 100

Slide 100 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: Xk+1 = Xk − γ∇V (Xk ) + 1 p 1 Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 p 2γZk+1 Wp (πγ , π) ≤ C γ 2p γ Generalization of [Brosse et al. ’19, Renaud et al. ’24] to weakly convex functions √ • Chain Yk+1 = Yk − γbγ (Yk ) + 2γZk+1 , with bγ (y ) = ∇f (y − γ∇g γ (y )) + ∇g γ (y ) - Geometric ergodicity: Wp (pYk , p∞ ) ≤ Ar kγ 1 1 - Drift: Wp (p∞ , pγ ) ≤ B∥bγ − ∇(f + g γ )∥ℓp (π1 ) ≤ B̃γ p 2 γ 13/18

Slide 101

Slide 101 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: Xk+1 = Xk − γ∇V (Xk ) + 1 p 1 Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 p 2γZk+1 Wp (πγ , π) ≤ C γ 2p γ Generalization of [Brosse et al. ’19, Renaud et al. ’24] to weakly convex functions √ • Chain Yk+1 = Yk − γbγ (Yk ) + 2γZk+1 , with bγ (y ) = ∇f (y − γ∇g γ (y )) + ∇g γ (y ) - Geometric ergodicity: Wp (pYk , p∞ ) ≤ Ar kγ 1 1 - Drift: Wp (p∞ , pγ ) ≤ B∥bγ − ∇(f + g γ )∥ℓp (π1 ) ≤ B̃γ p 2 γ 1 - Discretization: Wp (pγ , µγ ) ≤ C γ 2p 13/18

Slide 102

Slide 102 text

Sketch of proof Drift error: For two Markov Chains Xki defined by Lipschitz and strongly convex at infinity drifts bi : p i i Xk+1 = Xki − γbi (Xki ) + 2γZk+1 Discretization error: Let π ∝ e −V and define with invariant laws πγi , we have of invariant law πγ , then ∃γ̄, s.t. ∀γ ≤ γ̄: Xk+1 = Xk − γ∇V (Xk ) + 1 p 1 Wp (πγ1 , πγ2 ) ≤ B∥b1 − b2 ∥ℓ (π1 ) 2 p 2γZk+1 Wp (πγ , π) ≤ C γ 2p γ Generalization of [Brosse et al. ’19, Renaud et al. ’24] to weakly convex functions √ • Chain Yk+1 = Yk − γbγ (Yk ) + 2γZk+1 , with bγ (y ) = ∇f (y − γ∇g γ (y )) + ∇g γ (y ) - Geometric ergodicity: Wp (pYk , p∞ ) ≤ Ar kγ 1 1 - Drift: Wp (p∞ , pγ ) ≤ B∥bγ − ∇(f + g γ )∥ℓp (π1 ) ≤ B̃γ p 2 γ 1 - Discretization: Wp (pγ , µγ ) ≤ C γ 2p −f −g γ ✗ pYk ∝ ̸∝ e −f −g ∼e 13/18

Slide 103

Slide 103 text

Additional results Convergence to the true posterior distribution (Using only Moreau envelope properties) γ • With π ∝ e −f −g , µγ ∝ e −f −g and νγ = Proxγg #µγ , for p ≥ 1, we have lim Wp (µγ , π) = 0 , lim Wp (νγ , π) = 0 γ→0 γ→0 14/18

Slide 104

Slide 104 text

Additional results Convergence to the true posterior distribution (Using only Moreau envelope properties) γ • With π ∝ e −f −g , µγ ∝ e −f −g and νγ = Proxγg #µγ , for p ≥ 1, we have lim Wp (µγ , π) = 0 , lim Wp (νγ , π) = 0 γ→0 γ→0 • If g is L-Lipschitz, there exists Ep ∈ R+ such that ∀γ ∈ [0, L22 ] 1 1 Wp (µγ , π) ≤ Ep (L2 γ) p , Wp (νγ , π) ≤ Ep (L2 γ) p + Lγ 14/18

Slide 105

Slide 105 text

Additional results Convergence to the true posterior distribution (Using only Moreau envelope properties) γ • With π ∝ e −f −g , µγ ∝ e −f −g and νγ = Proxγg #µγ , for p ≥ 1, we have lim Wp (µγ , π) = 0 , lim Wp (νγ , π) = 0 γ→0 γ→0 • If g is L-Lipschitz, there exists Ep ∈ R+ such that ∀γ ∈ [0, L22 ] 1 1 Wp (µγ , π) ≤ Ep (L2 γ) p , Wp (νγ , π) ≤ Ep (L2 γ) p + Lγ PnP-PSGLA Inexact PSGLA with Dγ ≈ Proxγg [Ehrhardt et al. ’24]   p X̂k+1 = Dγ X̂k − γ∇f (X̂k ) + 2γZk+1 ✓ Robustness to approximations of the proximal operator 1 1 C Wp (pX̂k , µγ ) ≤ Ar kγ + Bγ 2p + 1 ∥Dγ − Proxγg ∥ℓ2p 2 (µ̂γ ) γp 14/18

Slide 106

Slide 106 text

3. Numerical illustrations

Slide 107

Slide 107 text

Gaussian example 15/18

Slide 108

Slide 108 text

Image restoration (MMSE estimation) • DiffPIR [Zhu et al. ’23]: not convergent • PnP-ULA [Laumont et al. ’22]: slower • PnP-PSGLA with TV denoiser [Ehrhardt et al. ’24]: convex regularizer • PnP-PSGLA with firmly non-expansive DnCNN denoiser from [Pesquet et al. ’21] 16/18

Slide 109

Slide 109 text

Image restoration 17/18

Slide 110

Slide 110 text

Conclusion • PSGLA ✓ New stability results in the non-convex setting ✗ Tightness of the bounds 18/18

Slide 111

Slide 111 text

Conclusion • PSGLA ✓ New stability results in the non-convex setting ✗ Tightness of the bounds • PnP-PSGLA vs PnP-ULA ✓ Computational time reduced... ✗ ... but still prohibitive 18/18

Slide 112

Slide 112 text

Conclusion • PSGLA ✓ New stability results in the non-convex setting ✗ Tightness of the bounds • PnP-PSGLA vs PnP-ULA ✓ Computational time reduced... ✗ ... but still prohibitive → Inertial schemes [Falk et al. ’25] 18/18

Slide 113

Slide 113 text

Thank you

Slide 114

Slide 114 text

References • M. Renaud, V. de Bortoli, A. Leclaire, N. Papadakis. From stability of Langevin diffusion to convergence of proximal MCMC for non-log-concave sampling, NeurIPS 2025 • M. Renaud, A. Leclaire, N. Papadakis. On the Moreau envelope properties of weakly convex functions, arXiv 2509.13960, 2025

Slide 115

Slide 115 text

Image restoration: MMSE standard deviation

Slide 116

Slide 116 text

Comparisons Algorithm DiffPIR [Zhu et al. ’23] RED [Romano et al. ’17] RED [Romano et al. ’17] PnP [Venkatakrishnan et al. ’13] PnP [Venkatakrishnan et al. ’13] PnP-ULA [Laumont et al. ’22] PnP-PSGLA [Ehrhardt et al. ’24] PnP-PSGLA Denoiser GSDRUNet DnCNN GSDRUNet DnCNN GSDRUNet DnCNN TV DnCNN PSNR ↑ 29.99 30.49 29.26 30.50 30.52 27.89 29.24 30.81 SSIM ↑ 0.88 0.89 0.88 0.91 0.92 0.82 0.89 0.92 LPIPS ↓ 0.06 0.06 0.12 0.06 0.07 0.12 0.08 0.05 N↓ 20 500 500 500 500 100,000 1,000 10,000 time (s) ↓ 1 6 20 6 20 1,200 25 120