Learning to Efficiently Sample from Diffusion Probabilistic Models

Daniel Watson, Jonathan Ho, Mohammad Norouzi, William Chan

Introduction

Denoising Diffusion Probabilistic Models (DDPMs) (Sohl-Dickstein et al., 2015; Ho et al., 2020) have emerged as a powerful class of generative models, which model the data distribution through an iterative denoising process. DDPMs have been applied successfully to a variety of applications, including unconditional image generation (Song and Ermon, 2019; Ho et al., 2020; Song et al., 2021; Nichol and Dhariwal, 2021), shape generation (Cai et al., 2020), text-to-speech (Chen et al., 2021; Kong et al., 2020) and single image super-resolution (Saharia et al., 2021; Li et al., 2021).

DDPMs are easy to train, featuring a simple denoising objective (Ho et al., 2020) with noise schedules that successfully transfer across different models and datasets. This contrasts to Generative Adversarial Networks (GANs) (Goodfellow et al., 2014), which require an inner-outer loop optimization procedure that often entails instability and requires careful hyperparameter tuning. DDPMs also admit a simple non-autoregressive inference process; this contrasts to autoregressive models with often prohibitive computational costs on high dimensional data. The DDPM inference process starts with samples from the corresponding prior noise distribution (e.g., standard Gaussian), and iteratively denoises the samples under the fixed noise schedule. However, DDPMs often need hundreds-to-thousands of denoising steps (each involving a feedforward pass of a large neural network) to achieve strong results. While this process is still much faster than autoregressive models, this is still often computationally prohibitive, especially when modeling high dimensional data.

There has been much recent work focused on improving the sampling speed of DDPMs. WaveGrad (Chen et al., 2021) introduced a manually crafted schedule requiring only 6 refinement steps; however, this schedule seems to be only applicable to the vocoding task where there is a very strong conditioning signal. Denoising Diffusion Implicit Models (DDIMs) (Song et al., 2020a) accelerate sampling from pre-trained DDPMs by relying on a family of non-Markovian processes. They accelerate the generative process through taking multiple steps in the diffusion process. However, DDIMs sacrifice the ability to compute log-likelihoods. Nichol and Dhariwal (2021) also explored the use of ancestral sampling with a subsequence of the original denoising steps, trying both a uniform stride and other hand-crafted strides. San-Roman et al. (2021) improve few-step sampling further by training a separate model after training a DDPM to estimate the level of noise, and modifying inference to dynamically adjust the noise schedule at every step to match the predicted noise level.

All these fast-sampling techniques rely on a key property of DDPMs – there is a decoupling between the training and inference schedule. The training schedule need not be the same as the inference schedule, e.g., a diffusion model trained to use 1000 steps may actually use only 10 steps during inference. This decoupling characteristic is typically not found in other generative models. In past work, the choice of inference schedule was often considered a hyperpameter selection problem, and often selected via intuition or extensive hyperparmeter exploration (Chen et al., 2021). In this work, we view the choice of inference schedule path as an independent optimization problem, wherein we attempt to learn the best schedule. Our approach relies on a dynamic programming algorithm, where given a fixed budget of KK refinement steps and a pre-trained DDPM, we find the set of timesteps that maximizes the corresponding evidence lower bound (ELBO). As an optimization objective, the ELBO has a key decomposability property: the total ELBO is the sum of individual KL terms, and for any two inference paths, if the timesteps (s,t)(s,t) contiguously occur in both, they share a common KL term, therefore admitting memoization.

Our main contributions are the following:

We introduce a dynamic programming algorithm that finds the optimal inference paths based on the ELBO for all possible computation budgets of KK refinement steps. The algorithm searches over T>KT>K timesteps, only requiring O(T)\mathcal{O}(T) neural network forward passes. It only needs to be applied once to a pre-trained DDPM, does not require training or retraining a DDPM, and is applicable to both time-discrete and time-continuous DDPMs.

We experiment with DDPM models from prior work. On both LsimpleL_{\textrm{simple}} CIFAR10 and LhybridL_{\textrm{hybrid}} ImageNet 64x64, we discover schedules which require only 32 refinement steps, yet sacrifice only 0.1 bits per dimension compared to their original counterparts with 1,000 and 4,000 steps, respectively.

Background on Denoising Diffusion Probabilistic Models

Denoising Diffusion Probabilistic Models (DDPMs) (Ho et al., 2020; Sohl-Dickstein et al., 2015) are defined in terms of a forward Markovian diffusion process qq and a learned reverse process pθp_{\theta}. The forward diffusion process gradually adds Gaussian noise to a data point x0{\bm{x}}_{0} through TT iterations,

where the scalar parameters α1:T\alpha_{1:T} determine the variance of the noise added at each diffusion step, subject to 0<αt<10<\alpha_{t}<1. The learned reverse process aims to model q(x0)q({\bm{x}}_{0}) by inverting the forward process, gradually removing noise from signal starting from pure Gaussian noise xT{\bm{x}}_{T},

The parameters of the reverse process can be optimized by maximizing the following variational lower bound on the training set,

Two notable properties of Gaussian diffusion process that help formulate DDPMs tractably and efficiently include:

Given the marginal distribution of xt{\bm{x}}_{t} given x0{\bm{x}}_{0} in (7), one can sample from the q(xt∣x0)q({\bm{x}}_{t}\mid{\bm{x}}_{0}) independently for different tt and perform SGD on a randomly chosen KL term in (6). Furthermore, given that the posterior distribution of xt−1{\bm{x}}_{t-1} given xt{\bm{x}}_{t} and x0{\bm{x}}_{0} is Gaussian, one can compute each KL term in (6) between two Gaussians in closed form and avoid high variance Monte Carlo estimation.

Linking DDPMs to Continuous Time Affine Diffusion Processes

Before describing our approach to efficiently sampling from DDPMs, it is helpful to link DDPMs to continuous time affine diffusion processes, as it shows the compatibility of our approach to both time-discrete and time-continuous DDPMs. Let x0∼q(x0){\bm{x}}_{0}\sim q({\bm{x}}_{0}) denote a data point drawn from the empirical distribution of interest and let q(xt∣x0)q({\bm{x}}_{t}|{\bm{x}}_{0}) denote a stochastic process for t∈t\in defined through an affine diffusion process through the following stochastic differential equation (SDE):

where fsde,gsde:→f_{\textrm{sde}},g_{\textrm{sde}}:\to are integrable functions satisfying fsde(0)=1f_{\textrm{sde}}(0)=1 and gsde(0)=0g_{\textrm{sde}}(0)=0.

Following Särkkä and Solin (2019) (section 6.1), we can compute the exact marginals q(xt∣xs)q({\bm{x}}_{t}|{\bm{x}}_{s}) for any 0≤s<t≤10\leq s<t\leq 1. This differs from Ho et al. (2020), where their marginals are those of the discretized diffusion via Euler-Maruyama, where it is not possible to compute marginals outside the discretization since they are formulated as cumulative products. We get:

where ψ(t,s)=exp⁡(∫stf(u)du)\psi(t,s)=\exp\left(\int_{s}^{t}f(u)du\right). Since these integrals are difficult to work with, we instead propose to define the marginals directly as

where f,g:→f,g:\rightarrow are differentiable, monotonic functions satisfying f(0)=1,f(1)=0,g(0)=0,g(1)=1f(0)=1,f(1)=0,g(0)=0,g(1)=1. Then, by implicit differentiation it follows that the corresponding diffusion is

To complete our formulation, let fts=f(t)f(s)f_{ts}=\frac{f(t)}{f(s)} and gts=g(t)2−fts2g(s)2g_{ts}=\sqrt{g(t)^{2}-f_{ts}^{2}g(s)^{2}}. Then, it follows that for any 0<s<t≤10<s<t\leq 1 we have that

justifying the compatibility of our main approach with time-continuous DDPMs. We note that this reverse process is also mathematically equivalent to a reverse process based on a time-discrete DDPM derived from a subsequence of the original timesteps as done by Song et al. (2020a); Nichol and Dhariwal (2021).

For the case of s=0s=0 in the reverse process, we follow the parametrization of Ho et al. (2020) to obtain discretized log likelihoods and compare our log likelihoods fairly with prior work.

Learning to Efficiently Sample from DDPMs

We now introduce our dynamic programming (DP) approach. In general, after training a DDPM, there is a decoupling between training and inference schedules. One can use a different inference schedule compared to training. Additionally, we can optimize a loss or reward function with respect to the timesteps themselves (after the DDPM is trained). In this paper, we use the ELBO as our objective, however we note that it is possible to directly optimize the timesteps with other metrics.

In our work, we choose to optimize ELBO as our objective. We rely on one key property of ELBO, its decomposability. We first make a few observations. The DDPM models the transition probability pθ(xs∣xt)p_{\theta}(x_{s}\mid x_{t}), or the cost to move from xt→xsx_{t}\rightarrow x_{s}. Given a pretrained DDPM, one can construct any valid ELBO path through it as long as two properties hold:

The path starts at t=0t=0 and ends at t=1t=1.

The path is contiguously connected without breaks.

In other words, the ELBO is a sum of individual ELBO terms that are functions of contiguous timesteps (ti′,ti−1′)(t^{\prime}_{i},t^{\prime}_{i-1}). Now the question remains, given a fixed budget KK steps, what is the optimal ELBO path?

First, we observe that any two paths that share a (t,s)(t,s) transition will share a common L(t,s)L(t,s) term. We exploit this property in our dynamic programming algorithm. When given a grid of timesteps 0=t0<t1<...<tT−1<tT=10=t_{0}<t_{1}<...<t_{T-1}<t_{T}=1 with T≥KT\geq K, it is possible to efficiently find the exact optimum (i.e., finding {t1′,...,tK−1′}⊂{t1,...,tT−1}\{t^{\prime}_{1},...,t^{\prime}_{K-1}\}\subset\{t_{1},...,t_{T-1}\} with the best ELBO) by memoizing all the individual L(t,s)L(t,s) ELBO terms for s,t∈{t0,...,tT}s,t\in\{t_{0},...,t_{T}\} with s<ts<t. We can then solve the canonical least-cost-path problem on a directed graph where s→ts\to t are nodes and the edge connecting them has cost L(t,s)L(t,s).

For time-continuous DDPMs, the choice of grid (i.e., the t1,...,tT−1t_{1},...,t_{T-1}) can be arbitrary. For models trained with discrete timesteps, the grid must be a subset of (or the full) original steps used during training, unless the model was regularized during training with methods such as the sampling procedure proposed by Chen et al. (2021).

2 Dynamic Programming Algorithm

We now outline our methodology to solve the least-cost-path problem. Our solution is similar to Dijkstra’s algorithm, but it differs to the classical least-cost-path problem where the latter is typically used, as our problem has additional constraints: we restrict our search to paths of exactly K+1K+1 nodes, and the start and end nodes are fixed.

Let CC and DD be (K+1)×(T+1)(K+1)\times(T+1) matrices. C[k,t]C[k,t] will be the total cost of the least-cost-path of length kk from tt to 0. DD will be filled with the timesteps corresponding to such paths; i.e., D[k,t]D[k,t] will be the timestep ss immediately previous to tt for the optimal kk-step path (assuming tt is also part of such path).

We initialize C=0C=0 and all the other C[0,⋅]C[0,\cdot] to ∞\infty (the D[0,⋅]D[0,\cdot] are irrelevant, but for ease of index notation we keep them in this section). Then, for each kk from 1 to KK, we iteratively set, for each tt,

where L(t,s)L(t,s) is the cost to transition from tt to ss (see Equation 17). For all s≥ts\geq t, we set L(t,s)=∞L(t,s)=\infty (e.g., we only move backwards in the diffusion process). This procedure captures the shortest path cost in CC and the shortest path itself in DD.

We further observe that running the DP algorithm for each kk from 1 to TT (instead of KK), we can extract the optimal paths for all possible budgets KK. Algorithm 1 illustrates a vectorized version of the procedure we have outlined in this section, while Algorithm 2 shows how to explicitly extract the optimal paths from DD.

3 Efficient Memoization

Experiments

We apply our method on a wide variety of pre-trained DDPMs from prior work. This emphasizes the fact that our method is applicable to any pre-trained DDPM model. In particular, we rely the CIFAR10 model checkpoints released by Nichol and Dhariwal (2021) on both their LhybridL_{\textrm{hybrid}} and LvlbL_{\textrm{vlb}} objectives. We also showcase results on CIFAR10 (Krizhevsky et al., 2009) with the exact configuration used by Ho et al. (2020), which we denote as LsimpleL_{\textrm{simple}}, as well as LhybridL_{\textrm{hybrid}} on ImageNet 64x64 (Deng et al., 2009) following Nichol and Dhariwal (2021), training these last two models ourselves for 800K and 3M steps, respectively, but otherwise using the exact same configurations as the authors.

In our experiments, we always search over a grid that includes all the timesteps used to train the model, i.e., {t/T:t∈{1,...,T−1}}\{t/T:t\in\{1,...,T-1\}\}. For our CIFAR10 results, we computed the memoization tables with Monte Carlo estimates over the full training dataset, while on ImageNet 64x64 we limited the number of datapoints in the Monte Carlo estimates to 16,384 images on the training dataset.

For each pre-trained model, we compare the negative log likelihoods (estimated using the full heldout dataset) of the strides discovered by our dynamic programming algorithm against even and quadratic strides, following Song et al. (2020a). We find that our dynamic programming algorithm discovers strides resulting in much better log likelihoods than the hand-crafted strides used in prior work, particularly in the few-step regime. We provide a visualization of the log likelihood curves as a function of computation budget in Figure 1 for LsimpleL_{\textrm{simple}} CIFAR10 and LhybridL_{\textrm{hybrid}} ImageNet 64x64 (Deng et al., 2009), a full list of the scores in the few-step regime in Table 1, and a visualization of the discovered steps themselves in Figure 2.

We further evaluate our discovered strides by reporting FID scores (Heusel et al., 2017) on 50,000 model samples against the same number of samples from the training dataset, as is standard in the literature. We find that, although our strides are yield much better log likelihoods, such optimization does not necessarily translate to also improving the FID scores. Results are included in Figure 3. This weakened correlation between log-likehoods and FID is consistent with observations in prior work (Ho et al., 2020; Nichol and Dhariwal, 2021).

2 Monte Carlo Ablation

To investigate the feasibility of our approach using minimal computation, we experimented with setting the number of Monte Carlo datapoints used to compute the dynamic programming table of negative log likelihood terms to 128 samples (i.e., easily fit into a single batch of GPU memory). We find that, for CIFAR10, the difference in log likelihoods is negligible, while on ImageNet 64x64 there is a visible yet slight improvement in negative log likelihood when filling the table with more samples. We hypothesize that this is due to the higher diversity of ImageNet. Nevertheless, we highlight that our procedure can be applied very quickly (i.e., with just TT forward passes of a neural network when using a single batch, as opposed to a running average over batches), even for large models, to significantly improve log their likelihoods in the few-step regime.

Related Work

DDPMs (Ho et al., 2020) have recently shown results that are competitive with GANs (Goodfellow et al., 2014), and they can be traced back to the work of Sohl-Dickstein et al. (2015) as a restricted family of deep latent variable models. Dhariwal and Nichol (2021) have more recently shown that DDPMs can outperform GANs in FID scores (Heusel et al., 2017). Song and Ermon (2019) have also linked DDPMs to denoising score matching (Vincent et al., 2008, 2010), which is crucial to the continuous-time formulation (Song et al., 2021). This connection to score matching has been explored further by Song and Kingma (2021), where other score-matching techniques (e.g., sliced score matching, Song et al. (2020b)) have been shown to be valid DDPM objectives and DDPMs are linked to energy-based models. More recent work on the few-step regime of DDPMs (Song et al., 2020a; Chen et al., 2021; Nichol and Dhariwal, 2021; San-Roman et al., 2021; Kong and Ping, 2021; Jolicoeur-Martineau et al., 2021) has also guided our research efforts. DDPMs are also very closely related to variational autoencoders (Kingma and Welling, 2013), where more recent work has shown that, with many stochastic layers, they can also attain competitive negative log likelihoods in unconditional image generation (Child, 2020). Also very closely related to DDPMs, there has also been work on non-autoregressive modeling of text sequences that can be regarded as discrete-space DDPMs with a forward process that masks or remove tokens (Lee et al., 2018; Gu et al., 2019; Stern et al., 2019; Chan et al., 2020; Saharia et al., 2020). The UNet architecture (Ronneberger et al., 2015) has been key to the recent success of DDPMs, and as shown by Ho et al. (2020); Nichol and Dhariwal (2021), augmenting UNet with self-attention (Shaw et al., 2018) in scales where attention is computationally feasible has helped bring DDPMs closer to the current state-of-the-art autoregressive generative models (Child et al., 2019; Jun et al., 2020; Roy et al., 2021).

Conclusion and Discussion

By regarding the selection of the inference schedule as an optimization problem, we present a novel and efficient dynamic programming algorithm to discover the optimal inference schedule for a pre-trained DDPM. Our DP algorithm finds an optimal inference schedule based on the ELBO given a fixed computation budget. Our method need only be applied once to discover the schedule, and does not require training or re-training the DPPM. In the few-step regime, we discover schedules on LsimpleL_{\textrm{simple}} CIFAR10 and LhybridL_{\textrm{hybrid}} ImageNet 64x64 that require only 32 steps, yet sacrifice ≤0.1\leq 0.1 bits per dimension compared to state-of-the-art DDPMs using hundreds-to-thousands of refinement steps. Our approach only needs forward passes of the DDPM neural network to fill the dynamic programming table of L(t,s)L(t,s) terms, and we show that we can fill the dynamic programming table with just O(T)\mathcal{O}(T) forward passes. Moreover, we show that we can estimate the table using only 128 Monte Carlo samples, finding this to be sufficient even for datasets such as ImageNet with high diversity. Our method achieves strong likelihoods with very few refinement steps, outperforming prior work utilizing hand-crafted strides (Ho et al., 2020; Nichol and Dhariwal, 2021).

Despite very strong log-likelihood results, especially in the few step regime, we observe limitations to our method. There is a disconnect between log-likehoods and FID scores, where improvements in log-likelihoods do not necessarily translate to improvements in FID scores. This is consistent with prior work, showing that the correlation between log likelihood and FID can be mismatched (Ho et al., 2020; Nichol and Dhariwal, 2021). We hope our work will encourage future research exploiting our general framework of optimization post-training in DDPMs, potentially utilizing gradient-based optimization over not only the ELBO, but also other, non-decomposable metrics. We particularly note that other sampling steps such as MCMC corrector steps or alternative predictor steps (e.g., following the reverse SDE) (Song et al., 2021) can also be incorporated into computation budget, and general learning frameworks like reinforcement learning are well-suited to explore this space as well as non-differentiable learning signals.

References

Appendix A Appendix

From Equation 10, we get by implicit differentiation that

Similarly as above and also using the fact that ψ(t,s)=ψ(t,0)ψ(s,0)\psi(t,s)=\frac{\psi(t,0)}{\psi(s,0)},

A.2 Proof for Equations 13 and 14

From Equation 10 and ψ(t,s)=ψ(t,0)ψ(s,0)\psi(t,s)=\frac{\psi(t,0)}{\psi(s,0)} it is immediate that ftsf_{ts} is the mean of q(xt∣xs)q(x_{t}|x_{s}). To show that gts2g_{ts}^{2} is the variance of q(xt∣xs)q(x_{t}|x_{s}), Equation 10 implies that

The mean of q(xs∣xt,x0)q(x_{s}|x_{t},x_{0}) is given by the Gaussian conjugate prior formula (where all the distributions are conditioned on x0x_{0}). Let μ=ftsxs\mu=f_{ts}x_{s}, so we have a prior over μ\mu given by

Then it follows by the formula that μ∣xt,x0\mu|x_{t},x_{0} has variance