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 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 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 refinement steps. The algorithm searches over timesteps, only requiring 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 CIFAR10 and 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 and a learned reverse process . The forward diffusion process gradually adds Gaussian noise to a data point through iterations,
where the scalar parameters determine the variance of the noise added at each diffusion step, subject to . The learned reverse process aims to model by inverting the forward process, gradually removing noise from signal starting from pure Gaussian noise ,
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 given in (7), one can sample from the independently for different and perform SGD on a randomly chosen KL term in (6). Furthermore, given that the posterior distribution of given and 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 denote a data point drawn from the empirical distribution of interest and let denote a stochastic process for defined through an affine diffusion process through the following stochastic differential equation (SDE):
where are integrable functions satisfying and .
Following Särkkä and Solin (2019) (section 6.1), we can compute the exact marginals for any . 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 . Since these integrals are difficult to work with, we instead propose to define the marginals directly as
where are differentiable, monotonic functions satisfying . Then, by implicit differentiation it follows that the corresponding diffusion is
To complete our formulation, let and . Then, it follows that for any 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 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 , or the cost to move from . Given a pretrained DDPM, one can construct any valid ELBO path through it as long as two properties hold:
The path starts at and ends at .
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 . Now the question remains, given a fixed budget steps, what is the optimal ELBO path?
First, we observe that any two paths that share a transition will share a common term. We exploit this property in our dynamic programming algorithm. When given a grid of timesteps with , it is possible to efficiently find the exact optimum (i.e., finding with the best ELBO) by memoizing all the individual ELBO terms for with . We can then solve the canonical least-cost-path problem on a directed graph where are nodes and the edge connecting them has cost .
For time-continuous DDPMs, the choice of grid (i.e., the ) 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 nodes, and the start and end nodes are fixed.
Let and be matrices. will be the total cost of the least-cost-path of length from to 0. will be filled with the timesteps corresponding to such paths; i.e., will be the timestep immediately previous to for the optimal -step path (assuming is also part of such path).
We initialize and all the other to (the are irrelevant, but for ease of index notation we keep them in this section). Then, for each from 1 to , we iteratively set, for each ,
where is the cost to transition from to (see Equation 17). For all , we set (e.g., we only move backwards in the diffusion process). This procedure captures the shortest path cost in and the shortest path itself in .
We further observe that running the DP algorithm for each from 1 to (instead of ), we can extract the optimal paths for all possible budgets . 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 .
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 and 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 , as well as 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., . 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 CIFAR10 and 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 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 CIFAR10 and ImageNet 64x64 that require only 32 steps, yet sacrifice 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 terms, and we show that we can fill the dynamic programming table with just 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 ,
A.2 Proof for Equations 13 and 14
From Equation 10 and it is immediate that is the mean of . To show that is the variance of , Equation 10 implies that
The mean of is given by the Gaussian conjugate prior formula (where all the distributions are conditioned on ). Let , so we have a prior over given by
Then it follows by the formula that has variance