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 DALLE 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 (e.g., natural images) into pure noise, whereas the reverse process transforms pure noise into samples from , 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 -accurate (i.e., uniformly accurate), as opposed to -accurate (see, e.g., [De ̵+21]). This is particularly problematic because the score matching objective is an loss (see Section 2 for details), and there are empirical studies suggesting that in practice, the score estimate is not in fact -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 , which we make more quantitative in Section 3:
The score function of the forward process is -Lipschitz.
The data distribution 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 is at most , then with an appropriate choice of step size, the SGM outputs a measure which is -close in total variation (TV) distance to in 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 with polynomial complexity, even when 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 -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 -accurate score estimate. The works [BMR22, LLT22] instead analyze SGMs under the more realistic assumption of an -accurate score estimate. However, the bounds of [BMR22] suffer from the curse of dimensionality, whereas the bounds of [LLT22] require 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 . 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 .
The forward process has the interpretation of transforming samples from the data distribution into pure noise. From the well-developed theory of Markov diffusions, it is known that if denotes the law of the OU process at time , then exponentially fast in various divergences and metrics such as the -Wasserstein metric ; 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 , which is the aim of generative modeling. In general, suppose that we have an SDE of the form
where 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 and set
then the process 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 is the reversed Brownian motion.For ease of notation, we do not distinguish between the forward and the reverse Brownian motions. Here, is called the score function for . Since (and hence for ) 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 , consider minimizing the loss over a function class ,
where 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 is independent of and , 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 from , leading to the finite-sample problem
where are i.i.d. standard Gaussians independent of the data . Moreover, if we parameterize the score function as , then the empirical problem is equivalent to
which has the illuminating interpretation of predicting the added noise from the noised data .
Discretization and implementation.
We now discuss the final steps required to obtain an implementable algorithm. First, in the learning phase, given samples from (e.g., a database of natural images), we train a neural network on the empirical score matching objective (2.8), see [SE19]. Let be the step size of the discretization; we assume that we have obtained a score estimate of for each time , where .
In order to approximately implement the reverse SDE (2.5), we first replace the score function with the estimate . Then, for we freeze the value of this coefficient in the SDE at time . It yields the new SDE
Since this is a linear SDE, it can be integrated in closed form; in particular, conditionally on , the next iterate has an explicit Gaussian distribution.
There is one final detail: although the reverse SDE (2.5) should be started at , we do not have access to directly. Instead, taking advantage of the fact that , we instead initialize the algorithm at , i.e., from pure noise.
Let denote the law of the algorithm at time . The goal of this work is to bound , taking into account three sources of error: (1) the estimation of the score function; (2) the discretization of the SDE with step size ; and (3) the initialization of the algorithm at rather than at .
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 leads to smaller discretization error, thereby furnishing an algorithm with gradient complexity (as opposed to sampling based on the overdamped Langevin process, which has complexity ), 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 is the law of the forward process at time . 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 , we arrive at the algorithm
for . 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 .
For all , the score is -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 , and unlike [LLT22] we do not assume that 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 be the output of the DDPM algorithm (Section 2.1) at time , and suppose that the step size satisfies , where . Then, it holds that
To interpret this result, suppose that and . Choosing and , and hiding logarithmic factors,
In particular, in order to have , it suffices to have score error .
We remark that the iteration complexity of 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 is not well-defined and hence Assumption 1 fails to hold. Also, the bound in Theorem 2 has a term involving which is infinite if is not absolutely continuous w.r.t. . As pointed out by [De ̵22], in general we cannot obtain non-trivial guarantees for , because has full support and therefore under the manifold hypothesis. Nevertheless, we show that we can apply our results using an early stopping technique.
Namely, consider the law of the OU process at a time , initialized at . Then, we show in Lemma 20 that, if where , then satisfies Assumption 1 with , , and . By substituting by into the result of Theorem 2, we obtain Corollary 3 below.
Taking as the new target corresponds to stopping the algorithm early: instead of running the algorithm backward for a time , we run the algorithm backward for a time (note that should be a multiple of the step size ).
Suppose that is supported on the ball of radius . Let . Then, the output of DDPM is -close in TV to the distribution , which is -close in to , provided that the step size is chosen appropriately according to Theorem 2 and
Suppose that is supported on the ball of radius . Let . Then, the output of the DDPM algorithm satisfies , provided that the step size is chosen appropriately according to Theorem 2 and and .
Finally, if the output of DDPM at time is projected onto for an appropriate choice of , then we can also translate our guarantees to the standard metric, which we state as the following corollary.
Suppose that is supported on the ball of radius . Let , and let denote the output of DDPM at time projected onto for . Then, it holds that , provided that the step size is chosen appropriately according to Theorem 2, , and .
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 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 , the score is -Lipschitz.
If we ignore the dependence on and assume that the score estimate is sufficiently accurate, then the iteration complexity guarantee of Theorem 2 is . 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 be the output of the SGM algorithm based on the CLD (Section 2.2) at time , and suppose that the step size satisfies , where . Then, there is a universal constant 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 process, not the increments of both the and 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 be the output of the SGM algorithm based on the CLD (Section 2.2) at time , where the data distribution is the standard Gaussian , and the score estimate is exact (). Suppose that the step size satisfies . Then, for the path measures and 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 , which leads to an iteration complexity that scales linearly in the dimension . 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 of the SGM is close to 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 for DDPM blows up at , but the score for CLD is well-defined at , 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 rather than at .
The two main ways to study Markov diffusions is via the -Wasserstein distance , or via information divergences such as the KL divergence or the divergence. In order for the reverse process to be contractive in the distance, one typically needs some form of log-concavity assumption for the data distribution . For example, if (i.e., ), then for the reverse process (2.5) we have
For , the coefficient in front of is positive; this shows that for times near , the reverse process is actually expansive, rather than contractive. This poses an obstacle for an analysis in . Although it is possible to perform a 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 . 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 , there are two salient proof techniques. The first is the interpolation method of [VW19] (originally for KL divergence, but extended to divergence in [Che+21]), which is the method used in [LLT22]. The interpolation method writes down a differential inequality for , which is used to bound in terms of and an additional error term. Unfortunately, the analysis of [LLT22] required taking to be the divergence, for which the interpolation method is quite delicate. In particular, the error term is bounded using a log-Sobolev assumption on , 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 -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 , and .
The reverse process (2.5) is denoted , where .
The SGM algorithm (2.9) is denoted , and . Recall that we initialize at , the standard Gaussian measure.
The process is the same as , except that we initialize this process at rather than at . We write .
Conventions for Girsanov’s theorem.
The three measures we consider over path space are:
, under which has the law of the reverse process (2.5);
, under which has the law of the SGM algorithm initialized at (corresponding to the process defined above).
We also use the following notion from stochastic calculus [Le ̵16, Definition 4.6]:
A local martingale is a stochastic process s.t. there exists a sequence of nondecreasing stopping times s.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 . When discussing quantities which involve both position and velocity (e.g., the joint distribution of ), 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 is also a -martingale and the process
is a Brownian motion under , the probability distribution with density w.r.t. .
If the assumptions of Girsanov’s theorem are satisfied (i.e., the condition (5.1)), we can apply Girsanov’s theorem to and
where This tells us that under , there exists a Brownian motion s.t.
Recall that under we have a.s.
The equation above still holds -a.s. since (even if is no longer a -Brownian motion). Plugging (5.4) into (5.5) we have -a.s.,We still have under because the marginal at time of is equal to the marginal at time of . That is a consequence of the fact that is a (true) -martingale.
In other words, under , the distribution of is the SGM algorithm started at , i.e., . 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 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 and denote the measures on path space corresponding to the reverse process (2.5) and the SGM algorithm with -accurate score estimate initialized at . Assume that and . Then,
Then, we give the approximation argument to prove the inequality (5.10).
Bound on the discretization error. For , 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 , the process is the time reversal of the forward process , we can apply the moment bounds in Lemma 10 and the movement bound in Lemma 11 to obtain
Recall that under we have a.s.
The equation above still holds -a.s. since . Combining the last two equations we then obtain -a.s.,
and In other words, is the law of the solution of the SDE (5.24). At this stage we have the bound
with a.s. and . Note that the distribution of (resp. ) is (resp. ).
Noting that for every and using Lemma 12, we have a.s., uniformly over . Therefore, 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, as , uniformly over . Therefore, using [AGS05, Corollary 9.4.6], as . Therefore,
We conclude with Pinsker’s inequality (). ∎
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 denote the forward process (2.1).
(score function bound) For all ,
Along the OU process, we have , where is independent of . Hence,
This follows from the -smoothness of [[, see, e.g.,]Lemma 9]vempala2019ulaisoperimetry. We give a short proof for the sake of completeness.
If is the generator associated with , then
Suppose that Assumption 2 holds. Let denote the forward process (2.1). For with , if , then
We omit the proofs of the two next lemmas as they are straightforward.
Then, for every , uniformly over . In particular, uniformly over .
5 Proof of Corollary 5
Proof. [Proof of Corollary 5] For , let denote the projection onto . We want to prove that . We use the decomposition
For the first term, since and both have support contained in , we can upper bound the Wasserstein distance by the total variation distance. Namely, [Rol22, Lemma 9] implies that
where is from Corollary 3, yielding
Next, we take so that . Since is -Lipschitz, we have
where is from Corollary 3. Combining these bounds,
We now take , , and 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 and consider
If we write , then the forward process satisfies the linear SDE
Since , is always invertible. Moreover, from , one can work out that the spectrum of is
However, is not diagonalizable. The case of is special, as it corresponds to the case when the spectrum is , and it corresponds to the critically damped case. Following [DVK22], which advocated for setting , 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 and , 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 and denote the measures on path space corresponding to the reverse process (2.11) and the SGM algorithm with -accurate score estimate initialized at . Assume that and . Then,
Proof. For , 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 , , and it follows that and . Substituting this into Lemma 16, we deduce that if , then
where in the last step we used .
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 . Let denote the path measure for the algorithm, and let denote the path measure for the continuous-time process. After applying Girsanov’s theorem, we obtain
In this expression, note that depends only on the position coordinate. Since the process is smoother (as we do not add Brownian motion directly to ), the error is of size , which allows us to take step size . 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 and , the error now involves controlling , which is of size (the process 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 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 such that for all ,
Since and , then . 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 , if then
Let denote the subspace . Then, since
where is well-defined on , we have
Here, is the measure on such that
Note that since , then if we write , we have
Let denote a mode. We bound
For the first term, [DKR22, Proposition 2] yields
For the second term, since the mode satisfies , we have
After combining the bounds, we obtain the claimed estimate (6.8).
Next, we consider the case of general . We have
We can apply (6.8) with in place of , noting that for which is -smooth for , to get
Next, we prove the moment and movement bounds for the CLD.
Suppose that Assumptions 2 and 4 hold. Let denote the forward process (2.10).
(score function bound) For all ,
Next, the coupling argument of [Che+18] shows that the CLD converges exponentially fast in the Wasserstein metric associated to a twisted norm which is equivalent (up to universal constants) to the Euclidean norm . It implies the following result, see, e.g., [Che+18, Lemma 8]:
Suppose that Assumptions 2 holds. Let denote the forward process (2.10). For with , if ,
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 is stationary for the forward process (2.10), we have for all . In this proof, since the score estimate is perfect and , we simply denote the path measure for the algorithm as . From Girsanov’s theorem in the form of Corollary 14 and from , we have
To lower bound this quantity, we use the inequality to write, for
Using the fact that and for all , we can then bound
provided that . 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 ).
Proof. [Proof of Lemma 19] Similarly to the proof of Theorem 7 above, we note that
Hence, for a universal constant (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 . Also, by the Cauchy–Schwarz inequality, we can give a crude bound: writing ,
where, by standard estimates on the supremum of Brownian motion [[, see, e.g.,]Lemma 23]chewi2021optimal, the first factor is finite if (again using independence across the dimensions). For the second factor, if we split the sum according to and use Hölder’s inequality,
provided , 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 , which completes the proof. ∎
Conclusion
In this work, we provided the first convergence guarantees for SGMs which hold under realistic assumptions (namely, -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 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 with an empirical average over samples from . The trick is to use an integration by parts identity to reformulate the objective. Here, will denote any constant that does not depend on the optimization variable . Expanding the square,
We can rewrite the second term using integration by parts:
where . Substituting this in,
where and are independent, and .
Appendix B Regularization
Suppose that where , and let denote the law of the OU process at time , started at . Let be such that and set . Then,
For every , satisfies Assumption 1 with
For the OU process (2.1), we have , where is independent of . Hence, for ,
We now take to ensure that . Since , it suffices to take .
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