ShiftDDPMs: Exploring Conditional Diffusion Models by Shifting Diffusion Trajectories
Zijian Zhang, Zhou Zhao, Jun Yu, Qi Tian
Introduction and Motivation
Deep generative models such as Generative Adversarial Networks (GANs) (Goodfellow et al. 2014), Variational Autoencoders (VAEs) (Kingma and Welling 2013), autoregressive models (Van Oord, Kalchbrenner, and Kavukcuoglu 2016) and normalizing flows (Rezende and Mohamed 2015) have shown remarkable abilities to model complex data distributions and synthesize high-quality samples in various fields. Diffusion models (Sohl-Dickstein et al. 2015) are recently brought back into focus by denoising diffusion probabilistic models (DDPMs) (Ho, Jain, and Abbeel 2020), which exhibits competitive image synthesis results and has been applied in a wide range of data modalities.
Generally, DDPMs gradually disrupt images by adding noise through a fixed forward process and learn its reverse process to generate samples from noise in a denoising way. There are two main methods to achieve conditional DDPMs. One is to learn an estimator that can compute the similarity between conditions and noisy data and use it to guide pre-trained unconditional DDPMs to sample towards specified conditions (Dhariwal and Nichol 2021). Another is to train a conditional DDPM from scratch by incorporating conditions into the function approximator of the reverse process. Both methods try to fit their conditional reverse process to the reversal of fixed unconditional forward process. This brings up a question: Can we design a more effective forward process utilizing given conditions to form a new type of conditional DDPMs and benefit from it?
We investigate this question by exploring the mechanism of how conditional DDPMs achieve conditional sampling based on unconditional forward process, similar to that in PDAE (Zhang, Zhao, and Lin 2022). We conduct some experiments, shown in Figure 1. Concretely, we train an unconditional DDPM and a conditional one on MNIST (LeCun et al. 1998), respectively. The conditional one incorporates class labels (one-hot vector) into the function approximator of parameterized reverse process. The top two rows respectively show the latents sampled from for various and the samples generated by the unconditional DDPM starting from corresponding latents. Intuitively, the latents for smaller preserve more high-level information (such as class) of corresponding data, and they will be totally lost when is large enough. It means that the diffusion trajectories originating from different data will get entangled, and the latents will become indistinguishable when is large. We then divide the diffusion trajectories into three stages: early-stage (), critical-stage () and late-stage (). Then we design a mixed sampling procedure that employs unconditional sampling but switches to conditional sampling during the specified stage. Note that the unconditional and conditional reverse process can be connected because they are trained to approximate the same forward process so that they recognize the same pattern of latents. The bottom three rows show the samples generated by three different mixed sampling procedures, where each row only employs conditional sampling for the right stage. As we can see, only the samples conditioned on input labels during critical-stage match the input class labels.
These phenomena show that, for unconditional forward process, the key to achieve conditional sampling is to shift and separate the generative trajectories of different conditions during critical-stage. Besides, to some extent, the training and sampling during early and late stages are independent of conditions and leave the condition modeling and generation to the limited critical-stage. If we can utilize extra latent space and allocate an exclusive diffusion trajectory for each condition to make the trajectories of different conditions disentangled all the time, it will disperse condition modeling to all timesteps and may improve the learning capacity of model.
Recently, Grad-TTS (Popov et al. 2021) and PriorGrad (Lee et al. 2021) introduce conditional forward process with data-dependent priors for audio diffusion models and enable more efficient training than those with unconditional forward process. However, their differences and connections have not been discussed, and there has not been a comprehensive exploration of this kind of methods, especially for image diffusion models. In this work, we systematically study how to design controllable diffusion trajectories according to conditions and its effect for conditional diffusion models. Our main contributions contain:
We systemically introduce conditional forward process for diffusion models and provide a unified point of view on existing related approaches.
By shifting diffusion trajectories, ShiftDDPMs improve the utilization rate of latent space and the learning capacity of model.
We demonstrate the feasibility and effectiveness of ShiftDDPMs on various image synthesis tasks with extensive experiments.
Related Works
Diffusion models (Sohl-Dickstein et al. 2015; Ho, Jain, and Abbeel 2020) are an emerging family of generative models and have exhibited remarkable abilities to synthesize high-quality samples. Numerous studies (Song et al. 2020; Song, Meng, and Ermon 2020; Dhariwal and Nichol 2021; Liu et al. 2022) and applications (Chen et al. 2020; Saharia et al. 2022; Huang et al. 2022a, b; Ye et al. 2022, 2023) have further improved and expanded diffusion models. Among existing practices of conditional diffusion models, only Grad-TTS (Popov et al. 2021) and PriorGrad (Lee et al. 2021) involve conditions in forward process but, nonetheless, they are totally different methods. We will demonstrate their differences under the point of view of ShiftDDPMs.
ShiftDDPMs
DDPMs (Ho, Jain, and Abbeel 2020) employ a forward process that sequentially destroys data distribution into with Markov diffusion kernels defined by a fixed variance schedule :
which admits sampling from for any timestep in closed form:
Then a parameterized Markov chain is trained to fit the reversal of forward process, denoising an arbitrary Gaussian noise to a data sample:
Training is performed by maximizing the model log likelihood with some parameterization and simplication:
See Appendix A for full details of DDPMs.
Conditional Forward Process
We aim to shift the diffusion trajectories in some way related to conditions. An intuitive way is to directly rewrite the Gaussian distribution in Eq.(2) as:
Specifically, is the cumulative mean shift of diffusion trajectories at -th step, where is a shift coefficient schedule that decides the shift mode and is a function which we call shift predictor that maps conditions into the latent space. is a diagonal covariance matrix, where is some function similar to . Comparing the diffusion trajectories to water pipes, then is employed to change their directions and is employed to change their size in latent space. Note that both and can be fixed or trainable. In our experiments on image synthesis, trainable leads to complex training and sampling procedure, unstable training and poor results, so we fix like that in Eq.(2). For generalization, we still use in our derivations. For simplicity, we use following substitution:
where and . We will discuss how to choose and in later sections.
With Eq.(6), we can derive corresponding forward diffusion kernels (See proof in Appendix A):
where (i.e. ). Intuitively, our forward diffusion kernels introduce a small perturbation conditioned on to original ones shown in Eq.(1).
With Eq.(6) and Eq.(7), the posterior distributions of forward steps for can be derived from Bayes’ rule (See proof in Appendix A):
Parameterized Reverse Process
The reverse process starts at , which is an approximation of , and employs parameterized kernels to fit .
According to Eq.(6), can be represented as:
where . Then we take it into Eq.(8) and derive the posterior mean of forward steps:
where all things are available except . We can employ a model to predict . Note that there is no need to feed into because we have encoded it into condition-dependent trajectories (i.e., in ) so that the model does not need its guidance.
Further improvements come from another parameterization because in Eq.(9) is given by:
where the second term is available. Therefore we can employ a model to predict the first term for training. We find this parameterization achieves better performance than predicting directly. Then we can get the predicted posterior distributions parameterized by :
Training Objective
With our conditional forward process and corresponding reverse process, our training objective can be represented as (See proof in Appendix A):
where is some constant, , , , , and for . During training, we follow DDPMs (Ho, Jain, and Abbeel 2020) to adopt the simplified training objective by uniformly sampling between and and ignoring loss weight . Algorithm 1 and Algorithm 2 describe our training and sampling procedure. Note that and will be optimized along with if they are trainable.
Intuitive Interpretation
Assume that and , DDPMs employ to predict , while ShiftDDPMs employ to predict . They are trained to predict the same pattern of objective but with different input (i.e. ). Compared with DDPMs, ShiftDDPMs transfer input condition onto diffusion trajectories by shifting to , which allows conditional training and sampling without feeding into the network. For DDPMs, only the training and sampling during limited critical-stage plays a key role for condition modeling and generation, while ShiftDDPMs disperse it to all timesteps and improve the utilization rate of latent space, which may lead to a better performance.
Furthermore, if is trainable, it will be optimized to find an optimal shift in latent space to specialize the diffusion trajectories of different conditions and make them disentangle as much as possible. The term in Eq.(10) will amend the sampling trajectories in every step to ensure they can finally fall on the data manifold.
Next, we will show that the forward process of Grad-TTS (Popov et al. 2021) and PriorGrad (Lee et al. 2021) correspond to a special choice of , respectively.
Prior-Shift
Grad-TTS (Popov et al. 2021) proposes a score-based text-to-speech generative model with the prior mean predicted by text encoder and aligner. Specifically, it defines a forward process satisfying the following SDE:
where corresponds to of our system ( represents the parameterized text encoder and aligner, represents the input text). We show that match a discretization of Eq.(14) (See proof in Appendix A). For forward process, increases from to and leads to shift to as increases. For reverse process, we have:
where because the reverse process starts from and it needs to eliminate the cumulative shift of forward process. From the view of diffusion trajectories, Grad-TTS changes the ending point of trajectories, so we name the shift mode as Prior-Shift.
Note that Grad-TTS still takes as an additional input to the score estimator, but we have stated that it is unnecessary. However, doing this will get at least not worse results, but also introduces additional parameter and computation.
Data-Normalization
PriorGrad (Lee et al. 2021) employs a forward process as follows:
where . Obviously, satisfies Eq.(16). For forward process, it first normalizes by subtracting its corresponding prior mean and then trains a diffusion model on normalized with prior . For reverse process, we have:
Intuitively, the reverse process starts from and has no amendments all the time except the last step, where it adds prior mean to the output (denormalization). From the view of diffusion trajectories, PriorGrad resets the starting point of trajectories on the data manifold, so we name the shift mode as Data-Normalization.
Unlike Prior-Shift that disperses the cumulative shift to all points on the diffusion trajectories, Data-Normalization does not disentangle the diffusion trajectories so that it must feed into the network to guide sampling. However, by carefully designing , it can achieve the same precision with a simpler network and have a faster convergence rate under some constraints (Lee et al. 2021). Data-Normalization is more suitable for variance-sensitive data such as audio.
Quadratic-Shift
Except for Prior-Shift, we propose a shift mode to disentangle the diffusion trajectories of different conditions by making the concave trajectories shown in Figure 1 convex. In this case, we don’t change their starting or ending point, and becomes a middle point, where they first progress to it and then go away from it. Therefore should be similar to some quadratic function opening downwards with and . Empirically, we choose . We name the shift mode as Quadratic-Shift.
Experiments
In this section, we conduct several conditional image synthesis experiments with ShiftDDPMs. Note that we always set . Full implementation details of all experiments can be found in Appendix B.
We first verify the effectiveness of ShiftDDPMs with three shift modes on toy dataset MNIST (LeCun et al. 1998). We employ two fixed shift predictors ( and ) and a trainable one ( with parameters ), mapping a one-hot vector to a matrix. Specifically, takes evenly spaced numbers over $32\times 32\bm{E}_{2}(\cdot)\bm{E}_{\psi}(\cdot)$ employs stacked transposed convolution layers to compute the matrix.
Figure 2 presents the conditional MNIST samples for different shift modes with different shift predictors. As we can see, all models work for conditional generation, and the visualization of learned for Prior-Shift and Quadratic-Shift contain the general shape of corresponding class, which means that they learn specialized trajectories for different conditions. Data-Normalization must feed into the model so it may ignore the shift.
Despite the success of the fixed shift predictor on MNIST, we get poor sample results when modeling complex data distribution such as CIFAR-10. Therefore we will always employ trainable shift predictor with parameter in the following experiments.
Sample Quality
We further evaluate ShiftDDPMs on CIFAR-10 (Krizhevsky and Hinton 2009). For a fair comparison, we retrain a DDPM as baseline (our DDPM) and then use the same experimental settings and resources to train other models. We train a traditional conditional DDPM (cond. DDPM) by incorporating class labels into the function approximator of reverse process. Moreover, we train a time-dependent classifier (Sohl-Dickstein et al. 2015; Song et al. 2020; Dhariwal and Nichol 2021) on noisy images and use its gradients to guide (our DDPM) to sample towards specified class (cls. DDPM). For ShiftDDPMs, we train three models, including Prior-Shift, Data-Normalization, and Quadratic-Shift, all with trainable shift predictors. Furthermore, we employ another two models (cond. Prior-Shift and cond. Quadratic-Shift) by incorporating class labels into the reverse process of Prior-Shift and Quadratic-Shift, with the same method with (cond. DDPM). Figure 3 presents some conditional CIFAR-10 samples generated by Quadratic-Shift. Table 1 shows Inception Score, FID, negative log likelihood for these models.
As we can see, our retrained unconditional DDPM is slightly better than the original one with the help of improved settings. With the help of conditional knowledge, conditional DDPM outperforms unconditional DDPM. Classifer-guided DDPM has poor results because it is sensitive to the classifier. Data-Normalization has an unstable training process and poor results, which means that it is not suitable for image synthesis. Both Prior-Shift and Quadratic-Shift outperform conditional DDPM, which proves that conditional forward process can improve the learning capacity of ShiftDDPMs. Although incorporating class labels can slightly improve their performance, it also introduces additional computational and parameter complexity.
Adaption to DDIM for Fast Sampling
DDIMs (Song, Meng, and Ermon 2020) generalize the forward process of DDPMs to non-Markovian process with an equivalent objective for training, which enables us to employ an accelerated reverse process with pre-trained DDPMs. Fortunately, ShiftDDPMs can be adapted to ShiftDDIMs. Specifically, we can generate from via:
where (See proof in Appendix A).
Then we employ , which is an increasing sub-sequence of of length , for accelerated sampling. The corresponding variance become , where is a hyperparameter that we can directly control. Figure 4 and Table 2 presents the conditional CIFAR-10 samples generated by Quadratic-Shift mode and its FID with different sampling steps and . ShiftDDIMs can still keep competitive FID even though it only samples for steps.
Interpolation of Diffusion Trajectories
DDPMs (Ho, Jain, and Abbeel 2020) show that one can interpolate the latents of two source data, decode the interpolated latent by the reverse process and get a sample similar to the interpolation of two source data. Inspired by this phenomenon, we can try to interpolate the diffusion trajectories of different conditions, which is equivalent to interpolating between different , such as for two different conditions and . In theory, decide the direction of diffusion trajectories and the interpolated will take the median direction, which can lead the reverse process to generate the samples with the mixed features of and .
We verify this idea by conducting the experiments of attribute-to-image (Yan et al. 2016) on LFW dataset (Huang et al. 2008). Specifically, it requires us to generate facial images according to the input attributes. Each image () in LFW corresponds to a 73-dim real-valued vector (), where the value of each dimension represents the degree of some attribute such as male, beard and so on. We employ Quadratic-Shift with a trainable shift predictor to train on the training set and evaluate it on the test set. Figure 5 presents some samples, which shows that ShiftDDPMs can learn a meaningful shift (like a heatmap of the face), and the generated images are consistent with the ground truth in labeled face attributes. Figure 6 presents the interpolations generated by Quadratic-Shift. The interpolations smoothly transition from one side to the other, which verifies our assumptions about the disentangled diffusion trajectories.
Image Inpainting
Except for class-conditional image synthesis, we conduct some image-to-image synthesis experiments. Compared with enumerable class label, image space is almost infinite and it is a challenge to assign a unique trajectory for each instance. To prove the capacity of ShiftDDPMs, we conduct image inpainting experiments using Irregular Mask Dataset (Liu et al. 2018) with three image datasets: CelebA-HQ (Liu et al. 2015), LSUN-church (Yu et al. 2015) and Places2 (Zhou et al. 2017). We employ Quadratic-Shift mode and a UNet based architecture as a shift predictor, which takes as input the masked image and predicts the shift. Figure 7 presents some inpainting samples. As we can see, ShiftDDPMs predict a template of complete image based on the masked one, which guides the trajectory to generate consistent and diverse completions. To further evaluate ShiftDDPMs on image inpainting, we follow prior works (Yu et al. 2019; Liu et al. 2018; Zhang et al. 2020) by reporting FID on Places2 dataset. We choose several GAN-based models: Contextual Attention (Yu et al. 2018), EdgeConnect (Nazeri et al. 2019) and StructureFlow (Ren et al. 2019) as baselines. Besides, we take score-based inpainting method proposed in (Song et al. 2020) as another baseline. Table 3 presents the quantitative results, and ShiftDDPMs achieve competitive results comparable to prior GAN-based methods. In addition, ShiftDDPMs also outperform the score-based inpainting method, showing that the extra utilization of the latent space to some extent improves the learning capacity of diffusion models.
Text-to-Image
We conduct text-to-image (text2img) experiments on CUB dataset (Wah et al. 2011). We employ Quadratic-Shift mode and a network as shift predictor to generate shift from the pre-trained sentence embeddings. Figure 8 presents some generated samples. We can see that the shift predictor can predict a meaningful template according to text and guide the trajectory to generate text-consistent images. We choose several GAN-based models GAN-INT-CLS (Reed et al. 2016), StackGAN (Zhang et al. 2017), StackGAN++ (Zhang et al. 2018) and AttnGAN (Xu et al. 2018) as baselines. Besides, we take traditional conditional diffusion method as another baseline, which only incorporates sentence embeddings into the function approximator of parameterized reverse process. Table 4 presents some quantitative results, and ShiftDDPMs achieve competitive results comparable to prior GAN-based methods and traditional conditional diffusion model.
The choice of is flexible. For Prior-Shift, any schedules of monotonically increasing from to can be applied on Prior-Shift. We have tried with following three types : , and and they all work well. Furthermore, can also be piecewise:
One can also design other reasonable . We leave empirical investigations of as future work.
Conclusion
In this work, we propose a novel and flexible conditional diffusion model called ShiftDDPMs by introducing conditional forward process with controllable condition-dependent diffusion trajectories. We analyze the differences of existing related methods under the point of view of ShiftDDPMs and first apply them on image synthesis. With ShiftDDPMs, we can achieve a better performance and learn some interesting features in latent space. Extensive qualitative and quantitative experiments on image synthesis demonstrate the feasibility and effectiveness of ShiftDDPMs.
Acknowledgments
This work was supported in part by the National Natural Science Foundation of China (Grant No.62020106007, No.U21B2040, No.62222211 and No.202100023), Zhejiang Natural Science Foundation (LR19F020006), Zhejiang Electric Power Co., Ltd. Science and Technology Project No.5211YF220006 and Yiwise.
References
Appendix A
Derivation of our conditional forward diffusion kernels
According to Markovian property, for all . Therefore, we can assume that:
As we have known the marginal Gaussian for :
from (Bishop 2006) (2.115), we can derive that the marginal Gaussian for , i.e., is given by:
We further consider the case for from following facts:
With similar derivation based on (Bishop 2006) (2.113), we can get:
which matches Eq.(23). Therefore we set , i.e., to make Eq.(25) true for .
One can also verify this conclusion with the recurrence relation in Eq.(25) by the rule of the sum of normally distributed random variables.
Derivation of the posterior distributions of our conditional forward frocess
For all , we can derive by Bayes’ rule:
From (Bishop 2006) (2.116 and 2.117), we have that is Gaussian and
Derivation of the training objective
The training objective can be represented as:
Then we can derive the first term by Gaussian probability density function:
Then we can derive the second term by Gaussian KullbackLeibler divergence:
Then we can derive the third term by Gaussian KullbackLeibler divergence:
Combining the above derivations, we can get final training objective:
where is some constant, , , , , and for .
A discretization of Grad-TTS
Grad-TTS defines a forward process with following SDE:
where corresponds to of our notations. Consider a discretization of it:
where because for Wiener process when . With this recurrence relation, we can derive that:
where we can get for Grad-TTS.
Appendix B
Implementation Details
We use the same settings with ADM (Dhariwal and Nichol 2021), including network architecture, timesteps, variance schedule, dropout, learning rate and EMA. We set batch size to for CIFAR-10, for LFW and for the others. We use feature map resolutions for models and for the others.
To compute , we employ a linear layer and stacked transposed convolution layers to map conditions (one-hot vector or attribute vector) to three-channel feature maps for CIFAR-10 and LFW dataset. For image inpainting on CelebA-HQ, LSUN-church, Place2 datasets, we employ a U-Net architecture for pixel-to-pixel prediction. For text-to-image synthesis on CUB bird dataset, we employ a linear layer and stacked transposed convolution layers with attention mechanism to map the pre-trained word embeddings to three-channel feature maps.
For image inpainting, we use Irregular Mask Dataset collected by (Liu et al. 2018), which contains 55,116 irregular raw masks for training and 24,866 for testing. During training, for each image in the batch, we first randomly sample a mask from 55,116 training masks, then perform some random augmentations on the mask, finally we use it to mask the image and get our class center for training. So the training masks are different all the time. The mask is irregular and may be 100% hole due to augmentations. During testing, we use 12,000 test masks sampled and augmented from 24,866 raw testing masks. These 12,000 masks are categorized by hole size according to hole-to-image area ratios (0-20%, 20-40%, 40-60%).
The classifier for (cls. DDPM) employs the encoder half UNet to classify the noisy images. For the class-conditional function approximator, we use AdaGN same with that in ADM (Dhariwal and Nichol 2021).
We train all our models on eight Nvidia RTX 2080Ti GPUs.