Sampling is as easy as learning the score: theory for diffusion models with minimal data assumptions

Sitan Chen, Sinho Chewi, Jerry Li, Yuanzhi Li, Adil Salim, Anru R. Zhang

Introduction

Score-based generative models (SGMs) are a family of generative models which achieve state-of-the-art performance for generating audio and image data [Soh+15, HJA20, DN21, Kin+21, Son+21, Son+21a, VKK21]; see, e.g., the recent surveys [Cao+22, Cro+22, Yan+22]. One notable example of an SGM are denoising diffusion probabilistic models (DDPMs) [Soh+15, HJA20], which are a key component in large-scale generative models such as DALL⋅\cdotE 2 [Ram+22]. As the importance of SGMs continues to grow due to newfound applications in commercial domains, it is a pressing question of both practical and theoretical concern to understand the mathematical underpinnings which explain their startling empirical successes.

As we explain in more detail in Section 2, at their mathematical core, SGMs consist of two stochastic processes, which we call the forward process and the reverse process. The forward process transforms samples from a data distribution qq (e.g., natural images) into pure noise, whereas the reverse process transforms pure noise into samples from qq, hence performing generative modeling. Implementation of the reverse process requires estimation of the score function of the law of the forward process, which is typically accomplished by training neural networks on a score matching objective [Hyv05, Vin11, SE19].

Providing precise guarantees for estimation of the score function is difficult, as it requires an understanding of the non-convex training dynamics of neural network optimization that is currently out of reach. However, given the empirical success of neural networks on the score estimation task, a natural and important question is whether or not accurate score estimation implies that SGMs provably converge to the true data distribution in realistic settings. This is a surprisingly delicate question, as even with accurate score estimates, as we explain in Section 2.1, there are several other sources of error which could cause the SGM to fail to converge. Indeed, despite a flurry of recent work on this question [De ̵+21, BMR22, De ̵22, Liu+22, LLT22, Pid22], prior analyses fall short of answering this question, for (at least) one of three main reasons:

Super-polynomial convergence. The bounds obtained are not quantitative (e.g., [De ̵+21, Liu+22, Pid22]), or scale exponentially in the dimension and other problem parameters [BMR22, De ̵22], and hence are typically vacuous for the high-dimensional settings of interest in practice.

Strong assumptions on the data distribution. The bounds require strong assumptions on the true data distribution, such as a log-Sobelev inequality (LSI) (see, e.g., [LLT22]). While the LSI is slightly weaker than log-concavity, it ultimately precludes the presence of substantial non-convexity, which impedes the application of these results to complex and highly multi-modal real-world data distributions. Indeed, obtaining a polynomial-time convergence analysis for SGMs that holds for multi-modal distributions was posed as an open question in [LLT22].

Strong assumptions on the score estimation error. The bounds require that the score estimate is L∞L^{\infty}-accurate (i.e., uniformly accurate), as opposed to L2L^{2}-accurate (see, e.g., [De ̵+21]). This is particularly problematic because the score matching objective is an L2L^{2} loss (see Section 2 for details), and there are empirical studies suggesting that in practice, the score estimate is not in fact L∞L^{\infty}-accurate (e.g., [ZC23]). Intuitively, this is because we cannot expect that the score estimate we obtain in practice will be accurate in regions of space where the true density is very low, simply because we do not expect to see many (or indeed, any) samples from such regions.

Providing an analysis which goes beyond these limitations is a pressing first step towards theoretically understanding why SGMs actually work in practice.

The concurrent and independent work of [LLT23] also obtains similar guarantees to our Corollary 3.

1 Our contributions

In this work, we take a step towards bridging theory and practice by providing a convergence guarantee for SGMs, under realistic (in fact, quite minimal) assumptions, which scales polynomially in all relevant problem parameters. Namely, our main result (Theorem 2) only requires the following assumptions on the data distribution qq, which we make more quantitative in Section 3:

The score function of the forward process is LL-Lipschitz.

The data distribution qq has finite KL divergence w.r.t. the standard Gaussian.

We note that all of these assumptions are either standard or, in the case of A2, far weaker than what is needed in prior work. Crucially, unlike prior works, we do not assume log-concavity, an LSI, or dissipativity; hence, our assumptions cover arbitrarily non-log-concave data distributions. Our main result is summarized informally as follows.

Under assumptions A1-A3, and if the score estimation error in L2L^{2} is at most O~(ε)\widetilde{O}(\varepsilon), then with an appropriate choice of step size, the SGM outputs a measure which is ε\varepsilon-close in total variation (TV) distance to qq in O~(L2d/ε2)\widetilde{O}(L^{2}d/\varepsilon^{2}) iterations.

We remark that our iteration complexity is actually quite tight: in fact, this matches state-of-the-art discretization guarantees for the Langevin diffusion [VW19, Che+21].

We find Theorem 1 to be quite surprising, because it shows that SGMs can sample from the data distribution qq with polynomial complexity, even when qq is highly non-log-concave (a task that is usually intractable), provided that one has access to an accurate score estimator. This answers the open question of [LLT22] regarding whether or not SGMs can sample from multimodal distributions, e.g., mixtures of distributions with bounded log-Sobolev constant. In the context of neural networks, our result implies that so long as the neural network succeeds at the learning task, the remaining part of the SGM algorithm based on the diffusion model is principled, in that it admits a strong theoretical justification.

In general, learning the score function is also a difficult task. Nevertheless, our result opens the door to further investigations, such as: do score functions for real-life data have intrinsic (e.g., low-dimensional) structure which can be exploited by neural networks? A positive answer to this question, combined with our sampling result, would then provide an end-to-end guarantee for SGMs.

More generally, our result can be viewed as a black-box reduction of the task of sampling to the task of learning the score function of the forward process, at least for distributions satisfying our mild assumptions. As a simple consequence, existing computational hardness results for learning natural high-dimensional distributions like mixtures of Gaussians [DKS17, Bru+21, GVV22] and pushforwards of Gaussians by shallow ReLU networks [DV21, Che+22, CLL22] immediately imply hardness of score estimation for these distributions. To our knowledge this yields the first known information-computation gaps for this task.

Given an L2L^{2}-accurate score estimate, SGMs can sample from (essentially) any data distribution.

This constitutes a powerful theoretical justification for the use of SGMs in practice.

Critically damped Langevin diffusion (CLD).

Using our techniques, we also investigate the use of the critically damped Langevin diffusion (CLD) for SGMs, which was proposed in [DVK22]. Although numerical experiments and intuition from the log-concave sampling literature suggest that the CLD could potentially speed up sampling via SGMs, we provide theoretical evidence to the contrary: in Section 3.3, we conjecture that SGMs based on the CLD do not exhibit improved dimension dependence compared to the original DDPM algorithm.

2 Prior work

We now provide a more detailed comparison to prior work, in addition to the previous discussion above.

By now, there is a vast literature on providing precise complexity estimates for log-concave sampling; see, e.g., the book draft [Che22] for an exposition to recent developments. The proofs in this work build upon the techniques developed in this literature. However, our work addresses the significantly more challenging setting of non-log-concave sampling.

The work of [De ̵+21] provides guarantees for the diffusion Schrödinger bridge [Son+21a]. However, as previously mentioned their result is not quantitative, and they require an L∞L^{\infty}-accurate score estimate. The works [BMR22, LLT22] instead analyze SGMs under the more realistic assumption of an L2L^{2}-accurate score estimate. However, the bounds of [BMR22] suffer from the curse of dimensionality, whereas the bounds of [LLT22] require qq to satisfy an LSI.

The recent work of [De ̵22], motivated by the manifold hypothesis, considers a different pointwise assumption on the score estimation error which allows the error to blow up at time and at spatial ∞\infty. We discuss the manifold setting in more detail in Section 3.2. Unfortunately, the bounds of [De ̵22] also scale exponentially in problem parameters such as the manifold diameter.

After the first version of this work appeared online, we became aware of two concurrent and independent works [Liu+22, LLT23] which share similarities with our work. Namely, [Liu+22] uses a similar proof technique as our Theorem 2 (albeit without explicit quantitative bounds), whereas [LLT23] obtains similar guarantees to our Corollary 3 below. The follow-up work of [CLL23] further improves upon the results in this paper.

We also mention that the use of reversed SDEs for sampling is also implicit in the interpretation of the proximal sampler algorithm [LST21] given in [Che+22a], and the present work can be viewed as expanding upon the theory of [Che+22a] using a different forward channel (the OU process).

Background on SGMs

In this section, we provide a brief exposition to SGMs, following [Son+21a].

In denoising diffusion probabilistic modeling (DDPM), we start with a forward process, which is a stochastic differential equation (SDE). For clarity, we consider the simplest possible choice, which is the Ornstein–Uhlenbeck (OU) process

but in this work we stick with the choice g≡1g\equiv 1.

The forward process has the interpretation of transforming samples from the data distribution qq into pure noise. From the well-developed theory of Markov diffusions, it is known that if qt≔law⁡(Xt)q_{t}\coloneqq\operatorname{law}(X_{t}) denotes the law of the OU process at time tt, then qt→γdq_{t}\to\gamma^{d} exponentially fast in various divergences and metrics such as the 22-Wasserstein metric W2W_{2}; see [BGL14].

Reverse process.

If we reverse the forward process (2.1) in time, then we obtain a process that transforms noise into samples from qq, which is the aim of generative modeling. In general, suppose that we have an SDE of the form

where (σt)t≥0{(\sigma_{t})}_{t\geq 0} is a deterministic matrix-valued process. Then, under mild conditions on the process (e.g., [Föl85, Cat+22]), which are satisfied for all processes under consideration in this work, the reverse process also admits an SDE description. Namely, if we fix the terminal time T>0T>0 and set

then the process (Xˉt←)t∈[0,T]{(\bar{X}^{\leftarrow}_{t})}_{t\in[0,T]} satisfies the SDE

where the backwards drift satisfies the relation

Applying this to the forward process (2.1), we obtain the reverse process

where now (Bt)t∈[0,T]{(B_{t})}_{t\in[0,T]} is the reversed Brownian motion.For ease of notation, we do not distinguish between the forward and the reverse Brownian motions. Here, ∇ln⁡qt\nabla\ln q_{t} is called the score function for qtq_{t}. Since qq (and hence qtq_{t} for t≥0t\geq 0) is not explicitly known, in order to implement the reverse process the score function must be estimated on the basis of samples.

Score matching.

In order to estimate the score function ∇ln⁡qt\nabla\ln q_{t}, consider minimizing the L2(qt)L^{2}(q_{t}) loss over a function class F\mathscr{F},

where F\mathscr{F} could be, e.g., a class of neural networks. The idea of score matching, which goes back to [Hyv05, Vin11], is that after applying integration by parts for the Gaussian measure, the problem (2.6) is equivalent to the following problem:

where Zt∼normal⁡(0,Id)Z_{t}\sim\operatorname{\mathsf{normal}}(0,I_{d}) is independent of Xˉ0\bar{X}_{0} and Xˉt=exp⁡(−t) Xˉ0+1−exp⁡(−2t) Zt\bar{X}_{t}=\exp(-t)\,\bar{X}_{0}+\sqrt{1-\exp(-2t)}\,Z_{t}, in the sense that (2.6) and (2.7) share the same minimizers. We give a self-contained derivation in Appendix A for the sake of completeness. Unlike (2.6), however, the objective in (2.7) can be replaced with an empirical version and estimated on the basis of samples Xˉ0(1),…,Xˉ0(n)\bar{X}_{0}^{(1)},\dotsc,\bar{X}_{0}^{(n)} from qq, leading to the finite-sample problem

where (Zt(i))i∈[n]{(Z_{t}^{(i)})}_{i\in[n]} are i.i.d. standard Gaussians independent of the data (Xˉ0(i))i∈[n]{(\bar{X}_{0}^{(i)})}_{i\in[n]}. Moreover, if we parameterize the score function as st=−11−exp⁡(−2t) z^ts_{t}=-\frac{1}{\sqrt{1-\exp(-2t)}}\,\widehat{z}_{t}, then the empirical problem is equivalent to

which has the illuminating interpretation of predicting the added noise Zt(i)Z_{t}^{(i)} from the noised data Xˉt(i)\bar{X}_{t}^{(i)}.

Discretization and implementation.

We now discuss the final steps required to obtain an implementable algorithm. First, in the learning phase, given samples Xˉ0(1),…,Xˉ0(n)\bar{X}_{0}^{(1)},\dotsc,\bar{X}_{0}^{(n)} from qq (e.g., a database of natural images), we train a neural network on the empirical score matching objective (2.8), see [SE19]. Let h>0h>0 be the step size of the discretization; we assume that we have obtained a score estimate skhs_{kh} of ∇ln⁡qkh\nabla\ln q_{kh} for each time k=0,1,…,Nk=0,1,\dotsc,N, where T=NhT=Nh.

In order to approximately implement the reverse SDE (2.5), we first replace the score function ∇ln⁡qT−t\nabla\ln q_{T-t} with the estimate sT−ts_{T-t}. Then, for t∈[kh,(k+1)h]t\in[kh,(k+1)h] we freeze the value of this coefficient in the SDE at time khkh. It yields the new SDE

Since this is a linear SDE, it can be integrated in closed form; in particular, conditionally on Xkh←X^{\leftarrow}_{kh}, the next iterate X(k+1)h←X^{\leftarrow}_{(k+1)h} has an explicit Gaussian distribution.

There is one final detail: although the reverse SDE (2.5) should be started at qTq_{T}, we do not have access to qTq_{T} directly. Instead, taking advantage of the fact that qT≈γdq_{T}\approx\gamma^{d}, we instead initialize the algorithm at X0←∼γdX^{\leftarrow}_{0}\sim\gamma^{d}, i.e., from pure noise.

Let pt≔law⁡(Xt←)p_{t}\coloneqq\operatorname{law}(X^{\leftarrow}_{t}) denote the law of the algorithm at time tt. The goal of this work is to bound TV⁡(pT,q)\operatorname{\mathsf{TV}}(p_{T},q), taking into account three sources of error: (1) the estimation of the score function; (2) the discretization of the SDE with step size h>0h>0; and (3) the initialization of the algorithm at γd\gamma^{d} rather than at qTq_{T}.

2 Background on the critically damped Langevin diffusion (CLD)

The critically damped Langevin diffusion (CLD) is based on the forward process

More generally, the CLD (2.10) is an instance of what is referred to as the kinetic Langevin or the underdamped Langevin process in the sampling literature. In the context of log-concave sampling, the smoother paths of Xˉ\bar{X} leads to smaller discretization error, thereby furnishing an algorithm with O~(d/ε)\widetilde{O}(\sqrt{d}/\varepsilon) gradient complexity (as opposed to sampling based on the overdamped Langevin process, which has complexity O~(d/ε2)\widetilde{O}(d/\varepsilon^{2})), see [Che+18, SL19, DR20, Ma+21]. In the recent paper [DVK22], Dockhorn, Vahdat, and Kreis proposed to use the CLD as the basis for an SGM and they empirically observed improvements over DDPM.

Applying (2.4), the corresponding reverse process is

where qt≔law⁡(Xˉt,Vˉt)\boldsymbol{q}_{t}\coloneqq\operatorname{law}(\bar{X}_{t},\bar{V}_{t}) is the law of the forward process at time tt. Note that the gradient in the score function is only taken w.r.t. the velocity coordinate. Upon replacing the score function with an estimate s\boldsymbol{s}, we arrive at the algorithm

for t∈[kh,(k+1)h]t\in[kh,(k+1)h]. We provide further background on the CLD in Section 6.1.

Results

We now state our assumptions and our main results.

For DDPM, we make the following mild assumptions on the data distribution qq.

For all t≥0t\geq 0, the score ∇ln⁡qt\nabla\ln q_{t} is LL-Lipschitz.

Assumption 1 is standard and has been used in the prior works [BMR22, LLT22]. However, unlike [LLT22], we do not assume Lipschitzness of the score estimate. Moreover, unlike [De ̵+21, BMR22], we do not assume any convexity or dissipativity assumptions on the potential UU, and unlike [LLT22] we do not assume that qq satisfies a log-Sobolev inequality. Hence, our assumptions cover a wide range of highly non-log-concave data distributions. Our proof technique is fairly robust and even Assumption 1 could be relaxed (as well as other extensions, such as considering the time-changed forward process (2.2)), although we focus on the simplest setting in order to better illustrate the conceptual significance of our results.

We also assume a bound on the score estimation error.

This is the same assumption as in [LLT22], and as discussed in Section 2.1, it is a natural and realistic assumption in light of the derivation of the score matching objective.

Our main result for DDPM is the following theorem.

Suppose that Assumptions 1, 2, and 3 hold. Let pTp_{T} be the output of the DDPM algorithm (Section 2.1) at time TT, and suppose that the step size h≔T/Nh\coloneqq T/N satisfies h≲1/Lh\lesssim 1/L, where L≥1L\geq 1. Then, it holds that

To interpret this result, suppose that KL⁡(q∥γd)≤poly⁡(d)\operatorname{\mathsf{KL}}(q\mathbin{\|}\gamma^{d})\leq\operatorname{poly}(d) and m2≤d\mathfrak{m}_{2}\leq d. Choosing T≍log⁡(KL⁡(q∥γd)/ε)T\asymp\log(\operatorname{\mathsf{KL}}(q\mathbin{\|}\gamma^{d})/\varepsilon) and h≍ε2L2dh\asymp\frac{\varepsilon^{2}}{L^{2}d}, and hiding logarithmic factors,

In particular, in order to have TV⁡(pT,q)≤ε\operatorname{\mathsf{TV}}(p_{T},q)\leq\varepsilon, it suffices to have score error εscore≤O~(ε)\varepsilon_{\rm score}\leq\widetilde{O}(\varepsilon).

We remark that the iteration complexity of N=Θ~(L2dε2)N=\widetilde{\Theta}(\frac{L^{2}d}{\varepsilon^{2}}) matches state-of-the-art complexity bounds for the Langevin Monte Carlo (LMC) algorithm for sampling under a log-Sobolev inequality (LSI), see [VW19, Che+21]. This provides some evidence that our discretization bounds are of the correct order, at least with respect to the dimension and accuracy parameters, and without higher-order smoothness assumptions.

2 Consequences for arbitrary data distributions with bounded support

For this setting, our results do not apply directly because the score function of qq is not well-defined and hence Assumption 1 fails to hold. Also, the bound in Theorem 2 has a term involving KL⁡(q∥γd)\operatorname{\mathsf{KL}}(q\mathbin{\|}\gamma^{d}) which is infinite if qq is not absolutely continuous w.r.t. γd\gamma^{d}. As pointed out by [De ̵22], in general we cannot obtain non-trivial guarantees for TV⁡(pT,q)\operatorname{\mathsf{TV}}(p_{T},q), because pTp_{T} has full support and therefore TV⁡(pT,q)=1\operatorname{\mathsf{TV}}(p_{T},q)=1 under the manifold hypothesis. Nevertheless, we show that we can apply our results using an early stopping technique.

Namely, consider qtq_{t} the law of the OU process at a time t>0t>0, initialized at qq. Then, we show in Lemma 20 that, if t≍εW22/(d (R∨d))t\asymp\varepsilon_{W_{2}}^{2}/(\sqrt{d}\,(R\vee\sqrt{d})) where 0<εW2≪d0<\varepsilon_{W_{2}}\ll\sqrt{d}, then qtq_{t} satisfies Assumption 1 with L≲dR2 (R∨d)2/εW24L\lesssim dR^{2}\,{(R\vee\sqrt{d})}^{2}/\varepsilon_{W_{2}}^{4}, KL⁡(qt∥γd)≤poly⁡(R,d,1/ε)\operatorname{\mathsf{KL}}(q_{t}\mathbin{\|}\gamma^{d})\leq\operatorname{poly}(R,d,1/\varepsilon), and W2(qt,q)≤εW2W_{2}(q_{t},q)\leq\varepsilon_{W_{2}}. By substituting qq by qtq_{t} into the result of Theorem 2, we obtain Corollary 3 below.

Taking qtq_{t} as the new target corresponds to stopping the algorithm early: instead of running the algorithm backward for a time TT, we run the algorithm backward for a time T−tT-t (note that T−tT-t should be a multiple of the step size hh).

Suppose that qq is supported on the ball of radius R≥1R\geq 1. Let t≍εW22/(d (R∨d))t\asymp\varepsilon_{W_{2}}^{2}/(\sqrt{d}\,(R\vee\sqrt{d})). Then, the output pT−tp_{T-t} of DDPM is εTV\varepsilon_{\rm TV}-close in TV to the distribution qtq_{t}, which is εW2\varepsilon_{W_{2}}-close in W2W_{2} to qq, provided that the step size hh is chosen appropriately according to Theorem 2 and

Suppose that qq is supported on the ball of radius R≥1R\geq 1. Let t≍ε2/(d (R∨d))t\asymp\varepsilon^{2}/(\sqrt{d}\,(R\vee\sqrt{d})). Then, the output pT−tp_{T-t} of the DDPM algorithm satisfies dBL(pT−t,q)≤ε\mathsf{d}_{\rm BL}(p_{T-t},q)\leq\varepsilon, provided that the step size hh is chosen appropriately according to Theorem 2 and N=Θ~(d3R4 (R∨d)4/ε10)N=\widetilde{\Theta}(d^{3}R^{4}\,(R\vee\sqrt{d})^{4}/\varepsilon^{10}) and εscore≤O~(ε)\varepsilon_{\rm score}\leq\widetilde{O}(\varepsilon).

Finally, if the output pT−tp_{T-t} of DDPM at time T−tT-t is projected onto B(0,R0)\mathsf{B}(0,R_{0}) for an appropriate choice of R0R_{0}, then we can also translate our guarantees to the standard W2W_{2} metric, which we state as the following corollary.

Suppose that qq is supported on the ball of radius R≥1R\geq 1. Let t≍ε2/(d (R∨d))t\asymp\varepsilon^{2}/(\sqrt{d}\,(R\vee\sqrt{d})), and let pT−t,R0p_{T-t,R_{0}} denote the output of DDPM at time T−tT-t projected onto B(0,R0)\mathsf{B}(0,R_{0}) for R0=Θ~(R)R_{0}=\widetilde{\Theta}(R). Then, it holds that W2(pT−t,R0,q)≤εW_{2}(p_{T-t,R_{0}},q)\leq\varepsilon, provided that the step size hh is chosen appropriately according to Theorem 2, N=Θ~(d3R8 (R∨d)4/ε12)N=\widetilde{\Theta}(d^{3}R^{8}\,(R\vee\sqrt{d})^{4}/\varepsilon^{12}), and εscore≤O~(ε)\varepsilon_{\rm score}\leq\widetilde{O}(\varepsilon).

Note that the dependencies in the three corollaries above are polynomial in all of the relevant problem parameters. In particular, since the last corollary holds in the W2W_{2} metric, it is directly comparable to [De ̵22] and vastly improves upon the exponential dependencies therein.

3 Results for CLD

In order to state our results for score-based generative modeling based on the CLD, we must first modify Assumptions 1 and 3 accordingly.

For all t≥0t\geq 0, the score ∇vln⁡qt\nabla_{v}\ln\boldsymbol{q}_{t} is LL-Lipschitz.

If we ignore the dependence on LL and assume that the score estimate is sufficiently accurate, then the iteration complexity guarantee of Theorem 2 is N=Θ~(d/ε2)N=\widetilde{\Theta}(d/\varepsilon^{2}). On the other hand, recall from Section 2.2 that based on intuition from the literature on log-concave sampling and from empirical findings in [DVK22], we might expect that SGMs based on the CLD have a smaller iteration complexity than DDPM. We prove the following theorem.

Suppose that Assumptions 2, 4, and 5 hold. Let pT\boldsymbol{p}_{T} be the output of the SGM algorithm based on the CLD (Section 2.2) at time TT, and suppose that the step size h≔T/Nh\coloneqq T/N satisfies h≲1/Lh\lesssim 1/L, where L≥1L\geq 1. Then, there is a universal constant c>0c>0 such that

Note that the result of Theorem 6 is in fact no better than our guarantee for DDPM in Theorem 2. Although it is possible that this is an artefact of our analysis, we believe that it is in fact fundamental. As we discuss in Remark 15, from the form of the reverse process (2.11), the SGM based on CLD lacks a certain property (that the discretization error should only depend on the size of the increment of the XX process, not the increments of both the XX and VV processes) which is crucial for the improved dimension dependence of the CLD over the Langevin diffusion in log-concave sampling. Hence, in general, we conjecture that under our assumptions, SGMs based on the CLD do not achieve a better dimension dependence than DDPM.

Let pT\boldsymbol{p}_{T} be the output of the SGM algorithm based on the CLD (Section 2.2) at time TT, where the data distribution qq is the standard Gaussian γd\gamma^{d}, and the score estimate is exact (εscore=0\varepsilon_{\rm score}=0). Suppose that the step size hh satisfies h≤110h\leq\frac{1}{10}. Then, for the path measures PT\boldsymbol{P}_{T} and QT←\boldsymbol{Q}^{\leftarrow}_{T} of the algorithm and the continuous-time process (2.11) respectively (see Section 6 for details), it holds that

Theorem 7 shows that in order to make the KL divergence between the path measures small, we must take h≲1/dh\lesssim 1/d, which leads to an iteration complexity that scales linearly in the dimension dd. Theorem 7 is not a proof that SGMs based on the CLD cannot achieve better than linear dimension dependence, as it is possible that the output pT\boldsymbol{p}_{T} of the SGM is close to q⊗γdq\otimes\gamma^{d} even if the path measures are not close, but it rules out the possibility of obtaining a better dimension dependence via our Girsanov-based proof technique. We believe that it provides compelling evidence for our conjecture, i.e., that under our assumptions, the CLD does not improve the complexity of SGMs over DDPM.

We remark that in this section, we have only considered the error arising from discretization of the SDE. It is possible that the score function for the SGM with the CLD is easier to estimate than the score function for DDPM, providing a statistical benefit of using the CLD. Indeed, under the manifold hypothesis, the score ∇ln⁡qt\nabla\ln q_{t} for DDPM blows up at t=0t=0, but the score ∇vln⁡qt\nabla_{v}\ln\boldsymbol{q}_{t} for CLD is well-defined at t=0t=0, and hence may lead to improvements over DDPM. We do not investigate this question here and leave it as future work.

Technical overview

We now give a detailed technical overview for the proof for DDPM (Theorem 2). The proof for CLD (Theorem 6) follows along similar lines.

Recall that we must deal with three sources of error: (1) the estimation of the score function; (2) the discretization of the SDE; and (3) the initialization of the reverse process at γd\gamma^{d} rather than at qTq_{T}.

The two main ways to study Markov diffusions is via the 22-Wasserstein distance W2W_{2}, or via information divergences such as the KL divergence or the χ2\chi^{2} divergence. In order for the reverse process to be contractive in the W2W_{2} distance, one typically needs some form of log-concavity assumption for the data distribution qq. For example, if ∇ln⁡q(x)=−x/σ2\nabla\ln q(x)=-x/\sigma^{2} (i.e., q∼normal⁡(0,σ2Id)q\sim\operatorname{\mathsf{normal}}(0,\sigma^{2}I_{d})), then for the reverse process (2.5) we have

For σ2≫1\sigma^{2}\gg 1, the coefficient in front of XˉT←\bar{X}^{\leftarrow}_{T} is positive; this shows that for times near TT, the reverse process is actually expansive, rather than contractive. This poses an obstacle for an analysis in W2W_{2}. Although it is possible to perform a W2W_{2} analysis using a weaker condition, such as a dissipativity condition, it typically leads to exponential dependence on the problem parameters (e.g., [De ̵22]).

On the other hand, the situation is different for an information divergence d\mathsf{d}. By the data-processing inequality, we always have

This motivates studying the processes via information divergences. We remark that the convergence of reversed SDEs has been studied in the context of log-concave sampling in [Che+22a] for the proximal sampler algorithm [LST21], providing the intuition behind these observations.

Next, we consider the score estimation error (1) and the discretization error (2). In order to perform a discretization analysis in KL or χ2\chi^{2}, there are two salient proof techniques. The first is the interpolation method of [VW19] (originally for KL divergence, but extended to χ2\chi^{2} divergence in [Che+21]), which is the method used in [LLT22]. The interpolation method writes down a differential inequality for ∂td(qT−t,pt)\partial_{t}\mathsf{d}(q_{T-t},p_{t}), which is used to bound d(qT−(k+1)h,p(k+1)h)\mathsf{d}(q_{T-(k+1)h},p_{(k+1)h}) in terms of d(qT−kh,pkh)\mathsf{d}(q_{T-kh},p_{kh}) and an additional error term. Unfortunately, the analysis of [LLT22] required taking d\mathsf{d} to be the χ2\chi^{2} divergence, for which the interpolation method is quite delicate. In particular, the error term is bounded using a log-Sobolev assumption on qq, see [Che+21] for further discussion. Instead, we pursue the second approach, which is to apply Girsanov’s theorem from stochastic calculus and to instead bound the divergence between measures on path space; this turns out to be doable using standard techniques. This is because, as noted in [Che+21], the Girsanov approach is more flexible as it requires less stringent assumptions.After the first draft of this work was made available online, we became aware of the concurrent and independent work of [Liu+22] which also uses an approach based on Girsanov’s theorem.

To elaborate, the main difficulty of using the interpolation method with an L2L^{2}-accurate score estimate (Assumption 3) is that the score estimation error is controlled by assumption under the law of the true process (2.5), but the interpolation analysis requires a control of the score estimation error under the law of the algorithm (2.9). Consequently, the work of [LLT22] required an involved change of measure argument in order to relate the errors under the two processes. In contrast, the Girsanov approach allows us to directly work with the score estimation error under the true process (2.5).

The forward process (2.1) is denoted (Xˉt)t∈[0,T]{(\bar{X}_{t})}_{t\in[0,T]}, and Xˉt∼qt\bar{X}_{t}\sim q_{t}.

The reverse process (2.5) is denoted (Xˉt←)t∈[0,T]{(\bar{X}^{\leftarrow}_{t})}_{t\in[0,T]}, where Xˉt←≔XˉT−t∼qT−t\bar{X}^{\leftarrow}_{t}\coloneqq\bar{X}_{T-t}\sim q_{T-t}.

The SGM algorithm (2.9) is denoted (Xt←)t∈[0,T]{(X^{\leftarrow}_{t})}_{t\in[0,T]}, and Xt←∼ptX^{\leftarrow}_{t}\sim p_{t}. Recall that we initialize at p0=γdp_{0}=\gamma^{d}, the standard Gaussian measure.

The process (Xt←,qT)t∈[0,T]{(X^{\leftarrow,q_{T}}_{t})}_{t\in[0,T]} is the same as (Xt←)t∈[0,T]{(X^{\leftarrow}_{t})}_{t\in[0,T]}, except that we initialize this process at qTq_{T} rather than at γd\gamma^{d}. We write Xt←,qT∼ptqTX^{\leftarrow,q_{T}}_{t}\sim p^{q_{T}}_{t}.

Conventions for Girsanov’s theorem.

The three measures we consider over path space are:

QT←Q^{\leftarrow}_{T}, under which (Xt)t∈[0,T]{(X_{t})}_{t\in[0,T]} has the law of the reverse process (2.5);

PTqTP^{q_{T}}_{T}, under which (Xt)t∈[0,T]{(X_{t})}_{t\in[0,T]} has the law of the SGM algorithm initialized at qTq_{T} (corresponding to the process (Xt←,qT)t∈[0,T]{(X^{\leftarrow,q_{T}}_{t})}_{t\in[0,T]} defined above).

We also use the following notion from stochastic calculus [Le ̵16, Definition 4.6]:

A local martingale (Lt)t∈[0,T](L_{t})_{t\in[0,T]} is a stochastic process s.t. there exists a sequence of nondecreasing stopping times Tn→TT_{n}\to T s.t. Ln=(Lt∧Tn)t∈[0,T]L^{n}=(L_{t\wedge T_{n}})_{t\in[0,T]} is a martingale.

Other parameters.

Notation for CLD.

The notational conventions for the CLD are similar; however, we must also consider a velocity variable VV. When discussing quantities which involve both position and velocity (e.g., the joint distribution qt\boldsymbol{q}_{t} of (Xˉt,Vˉt)(\bar{X}_{t},\bar{V}_{t})), we typically use boldface fonts.

Proofs for DDPM

First, we recall a consequence of Girsanov’s theorem that can be obtained by combining Pages 136–139, Theorem 5.22, and Theorem 4.13 of [Le ̵16].

then E⁡(L)\operatorname{\mathcal{E}}(\mathcal{L}) is also a QQ-martingale and the process

is a Brownian motion under P≔E⁡(L)T QP\coloneqq\operatorname{\mathcal{E}}(\mathcal{L})_{T}\,Q, the probability distribution with density E⁡(L)T\operatorname{\mathcal{E}}(\mathcal{L})_{T} w.r.t. QQ.

If the assumptions of Girsanov’s theorem are satisfied (i.e., the condition (5.1)), we can apply Girsanov’s theorem to Q=QT←Q=Q^{\leftarrow}_{T} and

where t∈[kh,(k+1)h].t\in[kh,(k+1)h]. This tells us that under P=E⁡(L)T QT←P=\operatorname{\mathcal{E}}(\mathcal{L})_{T}\,Q^{\leftarrow}_{T}, there exists a Brownian motion (βt)t∈[0,T]{(\beta_{t})}_{t\in[0,T]} s.t.

Recall that under QT←Q^{\leftarrow}_{T} we have a.s.

The equation above still holds PP-a.s. since P≪QT←P\ll Q^{\leftarrow}_{T} (even if BB is no longer a PP-Brownian motion). Plugging (5.4) into (5.5) we have PP-a.s.,We still have X0∼qTX_{0}\sim q_{T} under PP because the marginal at time t=0t=0 of PP is equal to the marginal at time t=0t=0 of QT←Q^{\leftarrow}_{T}. That is a consequence of the fact that E⁡(L)\operatorname{\mathcal{E}}(\mathcal{L}) is a (true) QT←Q^{\leftarrow}_{T}-martingale.

In other words, under PP, the distribution of XX is the SGM algorithm started at qTq_{T}, i.e., P=PTqT=E⁡(L)T QT←P=P^{q_{T}}_{T}=\operatorname{\mathcal{E}}(\mathcal{L})_{T}\,Q^{\leftarrow}_{T}. Therefore,

The equality (5.7) allows us to bound the discrepancy between the SGM algorithm and the reverse process.

2 Checking the assumptions of Girsanov’s theorem and the Girsanov discretization argument

In most applications of Girsanov’s theorem in sampling, a sufficient condition for (5.1) to hold, known as Novikov’s condition, is satisfied. Here, Novikov’s condition writes

and if Novikov’s condition holds, we can apply Girsanov’s theorem directly. However, under Assumptions 1, 2, and 3 alone, Novikov’s condition need not hold. Indeed, in order to check Novikov’s condition, we would want X0X_{0} to have sub-Gaussian tails for instance.

Furthermore, we also could not check that the condition (5.1), which is weaker than Novikov’s condition, holds. Therefore, in the proof of the next Theorem, we use a approximation technique to show that

We then use a discretization argument based on stochastic calculus to further bound this quantity. The result is the following theorem.

Suppose that Assumptions 1, 2, and 3 hold. Let QT←Q^{\leftarrow}_{T} and PTqTP^{q_{T}}_{T} denote the measures on path space corresponding to the reverse process (2.5) and the SGM algorithm with L2L^{2}-accurate score estimate initialized at qTq_{T}. Assume that L≥1L\geq 1 and h≲1/Lh\lesssim 1/L. Then,

Then, we give the approximation argument to prove the inequality (5.10).

Bound on the discretization error. For t∈[kh,(k+1)h]t\in[kh,(k+1)h], we can decompose

where the second term above is absorbed into the third term of the decomposition (5.16). Hence,

Using the fact that under QT←Q^{\leftarrow}_{T}, the process (Xt)t∈[0,T]{(X_{t})}_{t\in[0,T]} is the time reversal of the forward process (Xˉt)t∈[0,T]{(\bar{X}_{t})}_{t\in[0,T]}, we can apply the moment bounds in Lemma 10 and the movement bound in Lemma 11 to obtain

Recall that under QT←Q^{\leftarrow}_{T} we have a.s.

The equation above still holds PnP^{n}-a.s. since Pn≪QT←P^{n}\ll Q^{\leftarrow}_{T}. Combining the last two equations we then obtain PnP^{n}-a.s.,

and X0∼qT.X_{0}\sim q_{T}. In other words, PnP^{n} is the law of the solution of the SDE (5.24). At this stage we have the bound

with X0=X0nX_{0}=X^{n}_{0} a.s. and X0∼qTX_{0}\sim q_{T}. Note that the distribution of XnX^{n} (resp. XX) is PnP^{n} (resp. PTqTP^{q_{T}}_{T}).

Noting that Xtn=XtX^{n}_{t}=X_{t} for every t∈[0,Tn]t\in[0,T_{n}] and using Lemma 12, we have πε(Xn)→πε(X)\pi_{\varepsilon}(X^{n})\to\pi_{\varepsilon}(X) a.s., uniformly over [0,T][0,T]. Therefore, πε#Pn→πε#PTqT{\pi_{\varepsilon}}_{\#}P^{n}\to{\pi_{\varepsilon}}_{\#}P^{q_{T}}_{T} weakly. Using the lower semicontinuity of the KL divergence and the data-processing inequality [AGS05, Lemma 9.4.3 and Lemma 9.4.5], we obtain

Finally, using Lemma 13, πε(ω)→ω\pi_{\varepsilon}(\omega)\to\omega as ε→0\varepsilon\to 0, uniformly over [0,T][0,T]. Therefore, using [AGS05, Corollary 9.4.6], KL⁡((πε)#QT←∥(πε)#PTqT)→KL⁡(QT←∥PTqT)\operatorname{\mathsf{KL}}((\pi_{\varepsilon})_{\#}Q^{\leftarrow}_{T}\mathbin{\|}(\pi_{\varepsilon})_{\#}P^{q_{T}}_{T})\to\operatorname{\mathsf{KL}}(Q^{\leftarrow}_{T}\mathbin{\|}P^{q_{T}}_{T}) as ε↘0\varepsilon\searrow 0. Therefore,

We conclude with Pinsker’s inequality (TV⁡2≤KL⁡\operatorname{\mathsf{TV}}^{2}\leq\operatorname{\mathsf{KL}}). ∎

3 Proof of Theorem 2

Proof. [Proof of Theorem 2] We recall the notation from Section 4. By the data processing inequality,

Using the convergence of the OU process in KL divergence [[, see, e.g.,]Theorem 5.2.1]bakrygentilledoux2014 and applying Theorem 9 for the second term,

4 Auxiliary lemmas

In this section, we prove some auxiliary lemmas which are used in the proof of Theorem 2.

Suppose that Assumptions 1 and 2 hold. Let (Xˉt)t∈[0,T]{(\bar{X}_{t})}_{t\in[0,T]} denote the forward process (2.1).

(score function bound) For all t≥0t\geq 0,

Along the OU process, we have Xˉt=dexp⁡(−t) Xˉ0+1−exp⁡(−2t) ξ\bar{X}_{t}\overset{\mathsf{d}}{=}\exp(-t)\,\bar{X}_{0}+\sqrt{1-\exp(-2t)}\,\xi, where ξ∼normal⁡(0,Id)\xi\sim\operatorname{\mathsf{normal}}(0,I_{d}) is independent of Xˉ0\bar{X}_{0}. Hence,

This follows from the LL-smoothness of ln⁡qt\ln q_{t} [[, see, e.g.,]Lemma 9]vempala2019ulaisoperimetry. We give a short proof for the sake of completeness.

If Ltf≔Δf−⟨∇Ut,∇f⟩\mathscr{L}_{t}f\coloneqq\Delta f-\langle\nabla U_{t},\nabla f\rangle is the generator associated with qt∝exp⁡(−Ut)q_{t}\propto\exp(-U_{t}), then

Suppose that Assumption 2 holds. Let (Xˉt)t∈[0,T]{(\bar{X}_{t})}_{t\in[0,T]} denote the forward process (2.1). For 0≤s<t0\leq s<t with δ≔t−s\delta\coloneqq t-s, if δ≤1\delta\leq 1, then

We omit the proofs of the two next lemmas as they are straightforward.

Then, for every ε>0\varepsilon>0, fn→ff_{n}\to f uniformly over [0,T−ε][0,T-\varepsilon]. In particular, fn(⋅∧T−ε)→f(⋅∧T−ε)f_{n}(\cdot\wedge T-\varepsilon)\to f(\cdot\wedge T-\varepsilon) uniformly over [0,T][0,T].

5 Proof of Corollary 5

Proof. [Proof of Corollary 5] For R0>0R_{0}>0, let ΠR0\Pi_{R_{0}} denote the projection onto B(0,R0)\mathsf{B}(0,R_{0}). We want to prove that W2((ΠR0)#pT−t,q)≤εW_{2}((\Pi_{R_{0}})_{\#}p_{T-t},q)\leq\varepsilon. We use the decomposition

For the first term, since (ΠR0)#pT−t(\Pi_{R_{0}})_{\#}p_{T-t} and (ΠR0)#qt(\Pi_{R_{0}})_{\#}q_{t} both have support contained in B(0,R0)\mathsf{B}(0,R_{0}), we can upper bound the Wasserstein distance by the total variation distance. Namely, [Rol22, Lemma 9] implies that

where εTV\varepsilon_{\rm TV} is from Corollary 3, yielding

Next, we take R0≥RR_{0}\geq R so that (ΠR0)#q=q(\Pi_{R_{0}})_{\#}q=q. Since ΠR0\Pi_{R_{0}} is 11-Lipschitz, we have

where εW2\varepsilon_{W_{2}} is from Corollary 3. Combining these bounds,

We now take εW2=ε/3\varepsilon_{W_{2}}=\varepsilon/3, R0=Θ~(R)R_{0}=\widetilde{\Theta}(R), and εTV=Θ~(ε2/R2)\varepsilon_{\rm TV}=\widetilde{\Theta}(\varepsilon^{2}/R^{2}) to obtain the desired result. The iteration complexity follows from Corollary 3. ∎

Proofs for CLD

More generally, for the forward process we can introduce a friction parameter γ>0\gamma>0 and consider

If we write θˉt≔(Xˉt,Vˉt)\bar{\boldsymbol{\theta}}_{t}\coloneqq(\bar{X}_{t},\bar{V}_{t}), then the forward process satisfies the linear SDE

Since det⁡Aγ=1\det\boldsymbol{A}_{\gamma}=1, Aγ\boldsymbol{A}_{\gamma} is always invertible. Moreover, from tr⁡Aγ=−γ\operatorname{tr}\boldsymbol{A}_{\gamma}=-\gamma, one can work out that the spectrum of Aγ\boldsymbol{A}_{\gamma} is

However, Aγ\boldsymbol{A}_{\gamma} is not diagonalizable. The case of γ=2\gamma=2 is special, as it corresponds to the case when the spectrum is {−1}\{-1\}, and it corresponds to the critically damped case. Following [DVK22], which advocated for setting γ=2\gamma=2, we will also only consider the critically damped case. This also has the advantage of substantially simplifying the calculations.

2 Girsanov discretization argument

In order to apply Girsanov’s theorem, we introduce the path measures PTqT\boldsymbol{P}^{q_{T}}_{T} and QT←\boldsymbol{Q}^{\leftarrow}_{T}, under which

Applying Girsanov’s theorem, we have the following theorem.

Similarly to Appendix 5.2, even if Novikov’s condition does not hold, one can use an approximation to argue that the KL divergence is still upper bounded by the last expression. Since the argument follows along the same lines, we omit it for brevity.

Using this, we now aim to prove the following theorem.

Suppose that Assumptions 2, 4, and 5 hold. Let QT←\boldsymbol{Q}^{\leftarrow}_{T} and PTqT\boldsymbol{P}^{q_{T}}_{T} denote the measures on path space corresponding to the reverse process (2.11) and the SGM algorithm with L2L^{2}-accurate score estimate initialized at qT\boldsymbol{q}_{T}. Assume that L≥1L\geq 1 and h≲1/Lh\lesssim 1/L. Then,

Proof. For t∈[kh,(k+1)h]t\in[kh,(k+1)h], we can decompose

The change in the score function is bounded by Lemma 16, which generalizes [LLT22, Lemma C.12]. From the representation (6.1) of the solution to the CLD, we note that

In particular, since ∥A2∥op≲1\lVert\boldsymbol{A}_{2}\rVert_{\rm op}\lesssim 1, ∥A2−1∥op≲1\lVert\boldsymbol{A}_{2}^{-1}\rVert_{\rm op}\lesssim 1, and ∥Σ2∥op≲1\lVert\Sigma_{2}\rVert_{\rm op}\lesssim 1 it follows that ∥M0∥op=1+O(h)\lVert\boldsymbol{M}_{0}\rVert_{\rm op}=1+O(h) and ∥M1∥op=O(h)\lVert\boldsymbol{M}_{1}\rVert_{\rm op}=O(h). Substituting this into Lemma 16, we deduce that if h≲1/Lh\lesssim 1/L, then

where in the last step we used L≥1L\geq 1.

where the second term above is absorbed into the third term of the decomposition (6.7). Hence,

By applying the moment bounds in Lemma 17 together with Lemma 18 on the movement of the CLD process, we obtain

The proof is concluded via an approximation argument as in Section 5.2. ∎

Remark. We now pause to discuss why the discretization bound above does not improve upon the result for DDPM (Theorem 9). In the context of log-concave sampling, one instead considers the underdamped Langevin process

which is discretized to yield the algorithm

for t∈[kh,(k+1)h]t\in[kh,(k+1)h]. Let PT\boldsymbol{P}_{T} denote the path measure for the algorithm, and let QT\boldsymbol{Q}_{T} denote the path measure for the continuous-time process. After applying Girsanov’s theorem, we obtain

In this expression, note that ∇U\nabla U depends only on the position coordinate. Since the XX process is smoother (as we do not add Brownian motion directly to XX), the error ∥∇U(Xt)−∇U(Xkh)∥2\lVert\nabla U(X_{t})-\nabla U(X_{kh})\rVert^{2} is of size O(dh2)O(dh^{2}), which allows us to take step size h≲1/dh\lesssim 1/\sqrt{d}. This explains why the use of the underdamped Langevin diffusion leads to improved dimension dependence for log-concave sampling.

In contrast, consider the reverse process, in which

Since discretization of the reverse process involves the score function, which depends on both XX and VV, the error now involves controlling ∥Vt−Vkh∥2\lVert V_{t}-V_{kh}\rVert^{2}, which is of size O(dh)O(dh) (the process VV is not very smooth because it includes a Brownian motion component). Therefore, from the form of the reverse process, we may expect that SGMs based on the CLD do not improve upon the dimension dependence of DDPM.

In Section 6.5, we use this observation in order to prove a rigorous lower bound against discretization of SGMs based on the CLD.

3 Proof of Theorem 6

Proof. [Proof of Theorem 6] By the data processing inequality,

In [Ma+21], following the entropic hypocoercivity approach of [Vil09], Ma et al. consider a Lyapunov functional L\mathcal{L} which is equivalent to the sum of the KL divergence and the Fisher information,

which decays exponentially fast in time: there exists a universal constant c>0c>0 such that for all t≥0t\geq 0,

Since q0=q⊗γd\boldsymbol{q}_{0}=q\otimes\gamma^{d} and γ2d=γd⊗γd\boldsymbol{\gamma}^{2d}=\gamma^{d}\otimes\gamma^{d}, then L(q0∥γ2d)≲KL⁡(q∥γd)+FI⁡(q∥γd)\mathcal{L}(\boldsymbol{q}_{0}\mathbin{\|}\boldsymbol{\gamma}^{2d})\lesssim\operatorname{\mathsf{KL}}(q\mathbin{\|}\gamma^{d})+\operatorname{\mathsf{FI}}(q\mathbin{\|}\gamma^{d}). By Pinsker’s inequality and Theorem 15, we deduce that

4 Auxiliary lemmas

We begin with the perturbation lemma for the score function.

Proof. The proof follows along the lines of [LLT22, Lemma C.12]. First, we show that when M0=I2d\boldsymbol{M}_{0}=\boldsymbol{I}_{2d}, if L≤12 ∥M1∥opL\leq\frac{1}{2\,\lVert\boldsymbol{M}_{1}\rVert_{\rm op}} then

Let S\mathcal{S} denote the subspace S≔range⁡M1\mathcal{S}\coloneqq\operatorname{range}\boldsymbol{M}_{1}. Then, since

where M1−1\boldsymbol{M}_{1}^{-1} is well-defined on S\mathcal{S}, we have

Here, qθ\boldsymbol{q}_{\boldsymbol{\theta}} is the measure on θ+S\boldsymbol{\theta}+\mathcal{S} such that

Note that since L≤12 ∥M1∥opL\leq\frac{1}{2\,\lVert\boldsymbol{M}_{1}\rVert_{\rm op}}, then if we write qθ(θ′)∝exp⁡(−Hθ(θ′))\boldsymbol{q}_{\boldsymbol{\theta}}(\boldsymbol{\theta}^{\prime})\propto\exp(-\boldsymbol{H}_{\boldsymbol{\theta}}(\boldsymbol{\theta}^{\prime})), we have

Let θ⋆∈arg min⁡Hθ\boldsymbol{\theta}_{\star}\in\operatorname*{arg\,min}\boldsymbol{H}_{\boldsymbol{\theta}} denote a mode. We bound

For the first term, [DKR22, Proposition 2] yields

For the second term, since the mode satisfies ∇H(θ⋆)+M1−1 (θ⋆−θ)=0\nabla\boldsymbol{H}(\boldsymbol{\theta}_{\star})+\boldsymbol{M}_{1}^{-1}\,(\boldsymbol{\theta}_{\star}-\boldsymbol{\theta})=0, we have

After combining the bounds, we obtain the claimed estimate (6.8).

Next, we consider the case of general M0\boldsymbol{M}_{0}. We have

We can apply (6.8) with (M0)#q(\boldsymbol{M}_{0})_{\#}\boldsymbol{q} in place of q\boldsymbol{q}, noting that (M0)#q∝exp⁡(−H′)(\boldsymbol{M}_{0})_{\#}\boldsymbol{q}\propto\exp(-\boldsymbol{H}^{\prime}) for H′≔H∘M0\boldsymbol{H}^{\prime}\coloneqq\boldsymbol{H}\circ\boldsymbol{M}_{0} which is L′L^{\prime}-smooth for L′≔L ∥M0∥op2≲LL^{\prime}\coloneqq L\,\lVert\boldsymbol{M}_{0}\rVert_{\rm op}^{2}\lesssim L, to get

Next, we prove the moment and movement bounds for the CLD.

Suppose that Assumptions 2 and 4 hold. Let (Xˉt,Vˉt)t∈[0,T]{(\bar{X}_{t},\bar{V}_{t})}_{t\in[0,T]} denote the forward process (2.10).

(score function bound) For all t≥0t\geq 0,

Next, the coupling argument of [Che+18] shows that the CLD converges exponentially fast in the Wasserstein metric associated to a twisted norm ∣∣∣⋅∣∣∣\mathopen{|\mkern-1.5mu|\mkern-1.5mu|}\cdot\mathclose{|\mkern-1.5mu|\mkern-1.5mu|} which is equivalent (up to universal constants) to the Euclidean norm ∥⋅∥\lVert\cdot\rVert. It implies the following result, see, e.g., [Che+18, Lemma 8]:

Suppose that Assumptions 2 holds. Let (Xˉt,Vˉt)t∈[0,T]{(\bar{X}_{t},\bar{V}_{t})}_{t\in[0,T]} denote the forward process (2.10). For 0<s<t0<s<t with δ≔t−s\delta\coloneqq t-s, if δ≤1\delta\leq 1,

where we used the moment bound in Lemma 17. Next,

5 Lower bound against CLD

When proving upper bounds on the KL divergence, we can use the approximation argument described in Section 5.2 in order to invoke Girsanov’s theorem. However, when proving lower bounds on the KL divergence, this approach no longer works, so we check Novikov’s condition directly for the setting of Theorem 7.

Consider the setting of Theorem 7. Then, Novikov’s condition 6.2 holds.

We defer the proof of Lemma 19 to the end of this section. Admitting Lemma 19, we now prove Theorem 7.

Proof. [Proof of Theorem 7] Since q0=γd⊗γd=γ2d\boldsymbol{q}_{0}=\gamma^{d}\otimes\gamma^{d}=\boldsymbol{\gamma}^{2d} is stationary for the forward process (2.10), we have qt=γ2d\boldsymbol{q}_{t}=\boldsymbol{\gamma}^{2d} for all t≥0t\geq 0. In this proof, since the score estimate is perfect and qT=γ2d\boldsymbol{q}_{T}=\boldsymbol{\gamma}^{2d}, we simply denote the path measure for the algorithm as PT=PTqT\boldsymbol{P}_{T}=\boldsymbol{P}^{q_{T}}_{T}. From Girsanov’s theorem in the form of Corollary 14 and from sT−kh(x,v)=∇vln⁡qT−kh(x,v)=−v\boldsymbol{s}_{T-kh}(x,v)=\nabla_{v}\ln\boldsymbol{q}_{T-kh}(x,v)=-v, we have

To lower bound this quantity, we use the inequality ∥x+y∥2≥12 ∥x∥2−∥y∥2\lVert x+y\rVert^{2}\geq\frac{1}{2}\,\lVert x\rVert^{2}-\lVert y\rVert^{2} to write, for t∈[kh,(k+1)h]t\in[kh,(k+1)h]

Using the fact that Xˉs∼γd\bar{X}_{s}\sim\gamma^{d} and Vˉs∼γd\bar{V}_{s}\sim\gamma^{d} for all s∈[0,T]s\in[0,T], we can then bound

provided that h≤110h\leq\frac{1}{10}. Substituting this into (6.16),

This lower bound shows that the Girsanov discretization argument of Theorem 15 is essentially tight (except possibly the dependence on LL).

Proof. [Proof of Lemma 19] Similarly to the proof of Theorem 7 above, we note that

Hence, for a universal constant C>0C>0 (which may change from line to line)

By the Cauchy–Schwarz inequality, to prove that this expectation is finite, it suffices to consider the two terms in the exponential separately.

provided that h≲1/Th\lesssim 1/\sqrt{T}. Also, by the Cauchy–Schwarz inequality, we can give a crude bound: writing τ(t)=(exp⁡(2t)−1)/2\tau(t)=(\exp(2t)-1)/2,

where, by standard estimates on the supremum of Brownian motion [[, see, e.g.,]Lemma 23]chewi2021optimal, the first factor is finite if h≲1/Th\lesssim 1/\sqrt{T} (again using independence across the dimensions). For the second factor, if we split the sum according to exp⁡(−2t)≍2k\exp(-2t)\asymp 2^{k} and use Hölder’s inequality,

provided h≲1/Th\lesssim 1/T, where we again use [Che+21a, Lemma 23] and split across the coordinates. The Cauchy–Schwarz inequality then implies

For the second term, by independence of the increments,

By [Che+21a, Lemma 23], this quantity is finite if h≲1h\lesssim 1, which completes the proof. ∎

Conclusion

In this work, we provided the first convergence guarantees for SGMs which hold under realistic assumptions (namely, L2L^{2}-accurate score estimation and arbitrarily non-log-concave data distributions) and which scale polynomially in the problem parameters. Our results take a step towards explaining the remarkable empirical success of SGMs, at least under the assumption that the score function is learned with small L2L^{2} error.

The main limitation of this work is that we did not address the question of when the score function can be learned well. In general, studying the non-convex training dynamics of learning the score function via neural networks is challenging, but we believe that the resolution of this problem, even for simple learning tasks, would shed considerable light on SGMs. Together with the results in this paper, it would yield the first end-to-end guarantees for SGMs.

In another direction, and in light of the interpretation of our result as a reduction of the task of sampling to the task of score function estimation, we ask whether there are situations of interest in which it is easier to algorithmically learn the score function (not necessarily via a neural network) than it is to (directly) sample.

Acknowledgments. We thank Sébastien Bubeck, Yongxin Chen, Tarun Kathuria, Holden Lee, Ruoqi Shen, and Kevin Tian for helpful discussions. S. Chen was supported by NSF Award 2103300. S. Chewi was supported by the Department of Defense (DoD) through the National Defense Science & Engineering Graduate Fellowship (NDSEG) Program, as well as the NSF TRIPODS program (award DMS-2022448). A. Zhang was supported in part by NSF CAREER-2203741.

Appendix A Derivation of the score matching objective

In this section, we present a self-contained derivation of the score matching objective (2.7) for the reader’s convenience. See also [Hyv05, Vin11, SE19].

This objective cannot be evaluated, even if we replace the expectation over qtq_{t} with an empirical average over samples from qtq_{t}. The trick is to use an integration by parts identity to reformulate the objective. Here, CC will denote any constant that does not depend on the optimization variable sts_{t}. Expanding the square,

We can rewrite the second term using integration by parts:

where xt=exp⁡(−t) x0+1−exp⁡(−2t) ztx_{t}=\exp(-t)\,x_{0}+\sqrt{1-\exp(-2t)}\,z_{t}. Substituting this in,

where X0∼qX_{0}\sim q and Zt∼γdZ_{t}\sim\gamma^{d} are independent, and Xt≔exp⁡(−t) X0+1−exp⁡(−2t) ZtX_{t}\coloneqq\exp(-t)\,X_{0}+\sqrt{1-\exp(-2t)}\,Z_{t}.

Appendix B Regularization

Suppose that supp⁡q⊆B(0,R)\operatorname{supp}q\subseteq\mathsf{B}(0,R) where R≥1R\geq 1, and let qtq_{t} denote the law of the OU process at time tt, started at qq. Let ε>0\varepsilon>0 be such that ε≪d\varepsilon\ll\sqrt{d} and set t≍ε2/(d (R∨d))t\asymp\varepsilon^{2}/(\sqrt{d}\,(R\vee\sqrt{d})). Then,

For every t′≥tt^{\prime}\geq t, qt′q_{t^{\prime}} satisfies Assumption 1 with

For the OU process (2.1), we have Xˉt≔exp⁡(−t) Xˉ0+1−exp⁡(−2t) Z\bar{X}_{t}\coloneqq\exp(-t)\,\bar{X}_{0}+\sqrt{1-\exp(-2t)}\,Z, where Z∼normal⁡(0,Id)Z\sim\operatorname{\mathsf{normal}}(0,I_{d}) is independent of Xˉ0\bar{X}_{0}. Hence, for t≲1t\lesssim 1,

We now take t≲min⁡{ε/R,ε2/d}t\lesssim\min\{\varepsilon/R,\varepsilon^{2}/d\} to ensure that W22(q,qt)≤ε2W_{2}^{2}(q,q_{t})\leq\varepsilon^{2}. Since ε≪d\varepsilon\ll\sqrt{d}, it suffices to take t≍ε2/(d (R∨d))t\asymp\varepsilon^{2}/(\sqrt{d}\,(R\vee\sqrt{d})).

For this, we use the short-time regularization result in [OV01, Corollary 2], which implies that

Using [MS22, Lemma 4], along the OU process,

pages-1 rangepages14 rangepages24 rangepages41 rangepages16 rangepages31 rangepages13 rangepages15 rangepages38 rangepages12 rangepages15 rangepages33 rangepages37 rangepages8 rangepages12 rangepages12 rangepages15 rangepages12 rangepages-1 rangepages40 rangepages58 rangepages51 rangepages4 rangepages14 rangepages10 rangepages14 rangepages-1 rangepages-1 rangepages14 rangepages16 rangepages13

References