Deep Equilibrium Approaches to Diffusion Models

Ashwini Pokle, Zhengyang Geng, Zico Kolter

Introduction

Diffusion models have emerged as a promising class of generative models that can generate high quality images , outperforming GANs on perceptual quality metrics , and likelihood-based models on density estimation . One of the limitations of these models, however, is the fact that they require a long diffusion chain (many repeated applications of a denoising process), in order to generate high-fidelity samples. Several recent papers have focused on tackling this limitation, e.g., by shortening the length of diffusion process through an alternative parameterization , or through progressive distillation of a sampler with large diffusion chain into a smaller one . However, all of these methods still rely on a fundamentally sequential sampling process, imposing challenges on accelerating the sampling and for other applications like differentiating through the entire generation process.

In this paper, we propose an alternative approach that also begins to address such challenges from a different perspective. Specifically, we propose to model the generative process of a specific class of diffusion model, the denoising diffusion implicit model (DDIM) , as a deep equilibrium (DEQ) model . Deep equilibrium (DEQ) models are networks that aim to find the fixed point of the underlying system in the forward pass and differentiate implicitly through this fixed point in the backward pass. To apply DEQs to diffusion models, we first formulate this process as an equilibrium system consisting of all TT joint sampling steps, and then simultaneously solve the fixed point of all the TT steps to achieve sampling.

This approach has several benefits: First, the DEQ sampling process can be solved in parallel over multiple GPUs by batching the workload. This is particularly beneficial in the case of single image (i.e., batch-size-one) generation, where the serial sampling nature of diffusion models has inevitably made them unable to maximize GPU computation. Second, solving for the joint equilibria simultaneously leads to faster overall convergence as we have better estimates of the intermediate states in fewer steps. Specifically, the formulation naturally lends itself to an augmented diffusion chain that each state is generated according to all others. Third, the DEQ formulation allows us to leverage a faster differentiation through the chain. This enables us to much more effectively solve problems that require us to differentiate through the generative process, useful for tasks such as model inversion that seeks to find the noise that leads to a particular instance of an image.

We demonstrate the advantages of this formulation on two applications: single image generation and model inversion. Both single image generation and model inversion are widely applied in real world image manipulation tasks like image editing and restoration . On CIFAR-10 and CelebA , DEQ achieves up to 2×\times speedup over the sequential sampling processes of DDIM , while maintaining a comparable perceptual quality of images. In the model inversion, the loss converges much faster when trained with DEQs than with sequential sampling. Moreover, optimizing the sequential sampling process can be computationally expensive. Leveraging modern autograd packages is infeasible as they require storing the entire computational graph for all TT states. Some recent works like Nie et al. achieve this through use of SDE solvers. In contrast, with DEQs we can use the implicit differentiation of O(1)\mathcal{O}(1) memory complexity. Empirically, the initial hidden state recovered by DEQ more accurately regenerates the original image while capturing its finer details.

To summarize, our main contributions are as follows:

We formulate the generative process of an augmented type of DDIM as a deep equilibrium model that allows the use of black-box solvers to efficiently compute the fixed point and generate images.

The DEQ formulation parallelizes the sampling process of DDIM, and as a result, it can be run on multiple GPUs instead of a single GPU. This alternate sampling process converges faster compared to the original process.

We demonstrate the advantages of our formulation on single image generation and model inversion. We find that optimizing the sampling process via DEQs consistently outperforms naive sequential sampling.

We provide an easy way to extend this DEQ formulation to a more general family of diffusion models with stochastic generative process like DDPM .

Preliminaries

Denoising diffusion probabilistic models (DDPM) are generative models that can convert the data distribution to a simple distribution, (e.g., a standard Gaussian, N(0,I)\mathcal{N}(\textbf{0},{\bf{I}})), through a diffusion process. Specifically, given samples from a target distribution x0∼q(x0){\bf{x}}_{0}\sim q({\bf{x}}_{0}), the diffusion process is a Markov chain that adds Gaussian noises to the data to generate latent states x1,...,xT{\bf{x}}_{1},...,{\bf{x}}_{T} in the same sample space as x0{\bf{x}}_{0}. The inference distribution of diffusion process is given by:

To learn the parameters θ\theta that characterize a distribution pθ(x0)=∫pθ(x0:T)dx1:Tp_{\theta}({\bf{x}}_{0})=\int{p_{\theta}({\bf{x}}_{0:T})d{\bf{x}}_{1:T}} as an approximation of q(x0)q({\bf{x}}_{0}), a surrogate variational lower bound was proposed to train this model:

After training, samples can be generated by a reverse Markov chain, i.e., first sampling xT∼p(xT){\bf{x}}_{T}\sim p({\bf{x}}_{T}), and then repeatedly sampling xt−1{\bf{x}}_{t-1} till we reach x0{\bf{x}}_{0}.

As noted in , the length TT of a diffusion process is usually large (e.g., T=1000T=1000 ) as it contributes to a better approximation of Gaussian conditional distributions in the generative process. However, because of the large value of TT, sampling from diffusion models can be visibly slower compared to other deep generative models like GANs .

One feasible acceleration is to rewrite the forward process into a non-Markovian one that leads to a “shorter” and deterministic generative process, i.e., denoising diffusion implicit model (DDIM). DDIM can be trained similarly to DDPM, using the variational lower bound shown in Eq. 2. Essentially, DDIM constructs a nearly non-stochastic scheme that can quickly sample from the learned data distribution without introducing additional noises. Specifically, the scheme to generate a sample xt−1{\bf{x}}_{t-1} given xt{\bf{x}}_{t} is:

where α1,...,αT∈(0,1]\alpha_{1},...,\alpha_{T}\in(0,1], ϵt∼N(0,I){\bm{\epsilon}}_{t}\sim{\cal{N}}(\mathbf{0},{\bf{I}}), and ϵθ(t)(xt){\bm{\epsilon}}_{\theta}^{(t)}({\bf{x}}_{t}) is an estimator trained to predict the noise given a noisy state xt{\bf{x}}_{t}. For a variance schedule β1,…,βT\beta_{1},\ldots,\beta_{T}, we use the notation αt=∏s=1t(1−βs)\alpha_{t}=\prod_{s=1}^{t}(1-\beta_{s}). Different values of σt\sigma_{t} define different generative processes. When σt=(1−αt−1)/(1−αt)1−αt/αt−1\sigma_{t}=\sqrt{(1-\alpha_{t-1})/(1-\alpha_{t})}\sqrt{1-\alpha_{t}/\alpha_{t-1}} for all tt, the generative process represents a DDPM. Setting σt=0\sigma_{t}=0 for all tt gives rise to a DDIM, which results in a deterministic generating process except the initial sampling xT∼p(xT){\bf{x}}_{T}\sim p({\bf{x}}_{T}).

Deep Equilibrium Models

Deep equilibrium models are a recently-proposed class of deep networks that, in their forward pass, seek to find a fixed point of a single layer applied repeatedly to a hidden state. Specifically, consider a deep feedforward model with LL layers:

where x{\bf{x}} is the input injection, z[i]{\bf{z}}^{[i]} is the hidden state of ithi^{th} layer, and fθ[i]f_{\theta}^{[i]} is a layer that defines the feature transformation. Assuming the above model is weight-tied, i.e., fθ[i]=fθ,∀if_{\theta}^{[i]}=f_{\theta},\forall i, then in the limit of infinite depth, the output z[i]{\bf{z}}^{[i]} of this network converges to a fixed point z∗{\bf{z}}^{*}.

Inspired from the neural convergence phenomenon, Deep equilibrium (DEQ) models are proposed to directly compute this fixed point z∗{\bf{z}}^{*} as the output, i.e.,

The equilibrium state z∗{\bf{z}}^{*} can be solved by black-box solvers like Broyden’s method , or Anderson acceleration . To train this fixed-point system, Bai et al. leverage implicit differentiation to directly backpropagate through the equilibrium state z∗{\bf{z}}^{*} using O(1)\mathcal{O}(1) memory complexity. DEQ is known as a principled framework for characterizing convergence and energy minimization in deep learning. We leave a detailed discussion in Sec. 6.

A Deep Equilibrium Approach to DDIMs

In this section, we present the main modeling contribution of the paper, a formulation of diffusion processes under the DEQ framework. Although diffusion models may seem to be a natural fit for DEQ modeling (after all, we typically do not care about intermediate states in the denoising chain, but only the final clean image), there are several reasons why setting up the diffusion chain “naively” as a DEQ (i.e., making fθf_{\theta} be a single sampling step) does not ultimately lead to a functional algorithm. Most fundamentally, the diffusion process is not time-invariant (i.e., not “weight-tied” in the DEQ sense), and the final generated image is practically-speaking independent of the noise used to generate it (i.e., not truly based upon “input injection” either).

Thus, at a high level, our approach to building a DEQ version of the DDIM involves representing all the states x0:T{\bf{x}}_{0:T} simultaneously within the DEQ state. The advantage of this approach is that 1) we can exactly capture the typical diffusion inference chain; and 2) we can create a more expressive reverse process where the state xt{\bf{x}}_{t} is updated based upon all previous states xt+1:T{\bf{x}}_{t+1:T}, improving the inference process; 3) we can execute all steps of the inference chain in parallel rather than solely in sequence as is typically required in diffusion models; and 4) we can use common DEQ acceleration methods, such as the Anderson solver to find the fixed point, which makes the sampling process converge faster. A downside of this formulation is that we need to store all DEQ states simultaneously (i.e., only the images, not the intermediate network states).

The generative process of DDIM is given by:

This process also lets us generate a sample using a subset of latent states {xτ1,…,xτS}\{{\bf{x}}_{\tau_{1}},\ldots,{\bf{x}}_{\tau_{S}}\}, where {τ1,…,τS}⊆T\{\tau_{1},\ldots,\tau_{S}\}\subseteq T. While this helps in accelerating the overall generative process, there is a tradeoff between sampling quality and computational efficiency. As noted in Song et al. , larger TT values lead to lower FID scores of the generated images but need more compute time; smaller TT are faster to sample from, but the resulting images have worse FID scores.

Reformulating this sampling process as a DEQ addresses multiple concerns raised above. We can define a DEQ, with a sequence of latent states x1:T{\bf{x}}_{1:T} as its internal state, that simultaneously solves for the equilibrium points at all the timesteps. The global convergence of this process is upper bounded by TT steps, by definition. To derive the DEQ formulation of the generative process, first we rearrange the terms in Eq. (7):

Let c1(t)=1−αt−1−αt−1(1−αt)αtc_{1}^{(t)}=\sqrt{1-\alpha_{t-1}}-\sqrt{\dfrac{\alpha_{t-1}(1-\alpha_{t})}{\alpha_{t}}}. Then we can write

By induction, we can rewrite the above equation as:

This DEQ formulation has multiple benefits. Solving for all the equilibria simultaneously leads to a better estimation of the intermediate latent states xt{\bf{x}}_{t} in a fewer number of steps (i.e., ≤t\leq t steps for xt{\bf{x}}_{t}). This leads to faster convergence of the sampling process as the final sample x0{\bf{x}}_{0}, which is dependent on the latent states of all the previous time steps, has a better estimate of these intermediate latent states. Note that by the same reasoning, the intermediate latent states xt{\bf{x}}_{t} converge faster too. Thus, we can get images with perceptual quality comparable to DDIM in a significantly fewer number of steps. Of course, we also note that the computational requirements of each individual step has significantly increased, but this is at least largely offset by the fact that the steps can be executed as mini-batched in parallel over each state. Empirically, in fact, we often notice significant speedup using this approach on tasks like single image generation.

This DEQ formulation of DDIM can be extended to the stochastic generative processes of DDIM with η>0\eta>0, including that of DDPM (referred to as DEQ-sDDIM). The key idea is to sample noises for all the time steps along the sampling chain and treat this noise as an input injection to DEQ, in addition to xT{\bf{x}}_{T}.

where RootSolver(⋅\cdot) is any black-box fixed point solver, and ϵ1:T∼N(0,I){\bm{\epsilon}}_{1:T}\sim\mathcal{N}(\mathbf{0},{\bf{I}}) represents the input injected noises. We discuss this formulation in more detail in Appendix D.

Efficient Inversion of DDIM

One of the primary strengths of DEQs is their constant memory consumption, for both forward pass and backward pass, regardless of their ‘effective depth’. This leads to an interesting application of DEQs in inverting DDIMs that fully leverages this advantage along with the other benefits discussed in the previous section.

Given an arbitrary image x0∼D{\bf{x}}_{0}\sim\mathcal{D}, and a denoising diffusion model ϵθ(xt,t){\bm{\epsilon}}_{\theta}({\bf{x}}_{t},t) trained on a dataset D\mathcal{D}, model inversion seeks to determine the latent x^T∼N(0,I)\hat{{\bf{x}}}_{T}\sim{\cal{N}}(\mathbf{0},{\bf{I}}) that can generate an image x^0\hat{{\bf{x}}}_{0} identical to the original image x0{\bf{x}}_{0} through the generative process for DDIM described in Eq. (7). For an input image x0{\bf{x}}_{0}, and a generated image x^0\hat{{\bf{x}}}_{0}, this task needs to minimize the squared-Frobenius distance between these images:

2 Inverting DDIM: The Naive Approach

A relatively straightforward way to invert DDIM is to randomly sample xT∼N(0,I){\bf{x}}_{T}\sim{\cal{N}}(\mathbf{0},{\bf{I}}), and update it via gradient descent by first estimating x0{\bf{x}}_{0} using the generative process in Eq. (7) and backpropagating through this process after computing the loss objective in (14). The overall process has been summarized in Algorithm 1. This process has a large computational overhead. Every training epoch requires a sequential sampling for all TT timesteps. Optimizing through this generative process would require the creation of a large computational graph for storing relevant intermediate variables necessary for the backward pass. Sequential sampling further slows down the entire process.

3 Efficient Inversion of DDIM with DEQs

Alternatively, we can use the DEQ formulation to develop a much more efficient inversion method. We provide a high-level overview of this approach in Algorithm 2. We can apply implicit function theorem (IFT) to the fixed point, i.e., (12) to compute gradients of the loss L(x0,x0∗)\mathcal{L}({\bf{x}}_{0},{\bf{x}}_{0}^{*}) in (14) w.r.t. (⋅)(\cdot):

where (⋅)(\cdot) could be any of the latent states x1,...,xT{\bf{x}}_{1},...,{\bf{x}}_{T}, and J_{g_{\theta}}^{-1}\big{|}_{{\bf{x}}^{*}_{0:T}} is the inverse Jacobian of g(x0:T−1;xT)g({\bf{x}}_{0:T-1};{\bf{x}}_{T}) evaluated at x0:T∗x^{*}_{0:T}. Refer to for a detailed proof. Computing the inverse of Jacobian matrix can become computationally intractable, especially when the latent states xt{\bf{x}}_{t} are high dimensional. Further, prior works have reported growing instability of DEQs during training due to the ill-conditioning of Jacobian. Recent works suggest that we do not need an exact gradient to train DEQs. We can instead use an approximation to Eq. (15), i.e.,

where M{\bf{M}} is an approximation of J_{g_{\theta}}^{-1}\big{|}_{{\bf{x}}^{*}_{0:T}}. For example, show that setting M=I{\bf{M}}={\bf{I}}, i.e., 1-step gradient, works well. In this work, we follow Geng et al. to further add a damping factor to the 1-step gradient. The forward pass is given by:

The gradients for the backward pass can be computed through standard autograd packages. We provide the PyTorch-style pseudocode of our approach in the Appendix B. Using inexact gradients for the backward pass has several benefits: 1) It remarkably improves the training stability of DEQs; 2) Our backward pass consists of a single step and is ultra-cheap to compute. It reduces the total training time by a significant amount. It is easy to extend the strategy used in Algorithm 2 and use DEQs to invert DDIMs with stochastic generative process (referred to as DEQ-sDDIM). We provide the key steps of this approach in Algorithm 4.

Experiments

We consider four datasets that have images of different resolutions for our experiments: CIFAR10 (32×\times32) , CelebA (64×\times64) , LSUN Bedroom (256×\times256) and LSUN Outdoor Church (256×\times256) . For all the experiments, we use Anderson acceleration as the default fixed point solver. We use the pretrained denoising diffusion models from Ho et al. for CIFAR10, LSUN Bedroom, and LSUN Outdoor Church, and from Song et al. for CelebA. While training DEQs for model inversion, we use the 1-step gradient Eq. 18 to compute the backward pass. The damping factor τ\tau for 1-step gradient is set to 0.10.1. All the experiments have been performed on NVIDIA RTX A6000 GPUs. We provide additional experimental details in the Appendix A. While the primary focus in this section will be on the DDIM with a deterministic generative process i.e., η=0\eta=0, we also include a few key results on stochastic version of DDIM (DEQ-sDDIM) here. More extensive experiments can be found in Appendix D.

2 Sample quality of images generated with DEQ-DDIM

We verify that DEQ-DDIM can generate images of comparable quality to DDIM by reporting Fréchet Inception Distance (FID) in Table 1. For the forward pass of DEQ-DDIM, we run Anderson solver for a maximum of 15 steps for each image. We report FID scores on 50,000 images, and average time to generate an image (including GPU time) on 500 images. We note significant gains in wall-clock time on single-shot image generation with DEQ-DDIM on images with smaller resolutions. Specifically, DEQ-DDIM can generate images almost 2×\times faster than the sequential sampling of DDIM on CIFAR-10 (32×\times32) and CelebA (64×\times64). We note that these gains vanish on sequences of shorter lengths. This is because the number of fixed point solver iterations needed for convergence becomes comparable to the length of diffusion chain for small values of TT. Thus, lightweight updates performed on short diffusion chains for sequential sampling are faster compared to compute heavy updates in DEQ-DDIM.

We also report FID scores on DEQ-sDDIM for CIFAR10 in Table 2. We run Anderson solver for a maximum of 50 steps for each image. We observe that while DEQ-sDDIM is slower than DDIM, it always generates images with comparable or better FID scores. For higher levels of stochasticity i.e., for larger valued of η\eta, DEQ-sDDIM needs more Anderson solver iterations to converge to a fixed point, which increases image generation wall-clock time. We include additional results in Sec. D.2. Finally, we also find that on full-batch inference with larger batches, sequential sampling might outperform DEQ-DDIM, as the latter would have larger memory requirements in this case, i.e., processing smaller batches of size BB might be faster than processing larger batches of size BTBT.

3 Model Inversion of DDIM with DEQs

We report the minimum values of squared Frobenius norm between the recovered and target images averaged from 100 different runs in Table 3. We report results for DEQ with η=0\eta=0 (i.e., DEQ-DDIM) in this table, and additional results for η>0\eta>0 (i.e., DEQ-sDDIM) are reported in Figure 17. DEQ outperforms the baseline method on all the datasets by a significant margin. We also plot the training loss curves of DEQ-DDIM and the baseline in Figure 3. We observe that DEQ-DDIM converges faster and has much lower loss values than the baseline method induced by DDIM. We also visualize the images generated with the recovered latent states for DEQ-DDIM in Figure 4 and with DEQ-sDDIM in Figure 5. It is worth noting that images generated with DEQs capture more vivid details of the original images, like textures of foliage, crevices, and other finer details than the baseline. We include additional results of model inversion with DEQ-sDDIM on different datasets in Sec. D.3.

Related Work

Implicit deep learning is an emerging field that introduces structured methods to construct modern neural networks. Different from prior explicit counterparts defined by hierarchy or layer stacking, implicit models take advantage of dynamical systems , e.g., optimization , differential equation , or fixed-point system . For instance, Neural ODE describes a continuous time-dependent system, while Deep Equilibrium (DEQ) model , which is actually path-independent, is a new type of implicit models that outputs the equilibrium states of the underlying system, e.g., z∗{\bf{z}}^{*} from z∗=fθ(z∗,x){\bf{z}}^{*}=f_{\theta}({\bf{z}}^{*},{\bf{x}}) given the input x{\bf{x}}. This fixed-point system can be solved by black-box solvers , and further accelerated by the neural solver in the inference. An active topic is the stability of such a system as it will gradually deteriorate during training, albeit strong performances . DEQ has achieved SOTA results on a wide-range of tasks like language modeling , semantic segmentation , graph modeling , object detection , optical flow estimation , robustness , and generative models like normalizing flow , with theoretical guarantees .

Diffusion Models

Diffusion models , or score-based generative models , are newly developed generative models that utilize an iterative denoising process to progressively sample from a learned data distribution, which actually is the reverse of a forward diffusion process. They have demonstrated impressive fidelity for text-conditioned image generation and outperformed state-of-the-art GANs on ImageNet . Despite the superior practical results, diffusion models suffer from a plodding sampling speed, e.g., over hours to generate 50k CIFAR-sized images . To accelerate the sampling from diffusion models, researchers propose to skip a part of the sampling steps by reframing the reverse chain , or distill the trained diffusion model into a faster one . Plus, the forward and backward processes in diffusion models can be formulated as stochastic differential equations , bridging diffusion models with Neural ODEs in implicit deep learning. However, the community still lacks insights into the connection between DEQ and diffusion models, where we build our work to investigate this.

Model inversion

Model inversion gives insights into the latent space of a generative model, as an inability of a generative model to correctly reconstruct an image from its latent code is indicative of its inability to model all the attributes of image correctly. Further, the ability to manipulate the latent codes to edit high-level attributes of images finds applications in many tasks like semantic image manipulation , super resolution , in-painting , compressed sensing , etc. For generative models like GANs , inversion is non-trivial and requires alternate methods like learning the mapping from an image to the latent code , and optimizing the latent code through optimizers, e.g., both gradient-based and gradient-free . For diffusion models like DDPM , the generative process is stochastic, which can make model inversion very challenging. Many existing works based on diffusion models edit images or solve inverse problems without requiring full model inversion. Instead, they do so by utilizing existing understanding of diffusion models as presented in some recent works . Diffusion models have been widely applied to conditional image generation . Chung et al. propose a method to reduce the number of steps in reverse conditional diffusion process through better initialization, based on the idea of contraction theory of stochastic differential equations. Our proposed method is orthogonal to this work; we explicitly model DDIM as a joint, multivariate fixed point system and leverage black-box root solvers to solve for the fixed point and also allow for efficient differentiation.

Conclusion

We propose an approach to elegantly unify diffusion models and deep equilibrium (DEQ) models. We model the entire sampling chain of the denoising diffusion implicit model (DDIM) as a joint, multivariate (deep) equilibrium model. This setup replaces the traditional sequential sampling process with a parallel one, thereby enabling us to enjoy speedup obtained from multiple GPUs. Further, we can leverage inexact gradients to optimize the entire sampling chain quickly, which results in significant gains in model inversion. We demonstrate the benefits of this approach on 1) single-shot image generation, where we were able to obtain FID scores on par with or slightly better than those of DDIM; and 2) model inversion, where we achieved much faster convergence. We also propose an easy way to extend DEQ formulation for deterministic DDIM to its stochastic variants. It is possible to further speedup the sampling process by training a DEQ model to predict the noise at a particular timestep of the diffusion chain. We can jointly optimize the noise prediction network, and the latent variables of the diffusion chain, which we leave as future work.

Acknowledgements

Ashwini Pokle is supported by a grant from the Bosch Center for Artificial Intelligence.

References

Appendix A Experimental Details

In this section, we present detailed settings for all the experiments in Section 5.

We use exactly the same U-Net architecture for ϵθ(xt,t){\bm{\epsilon}}_{\theta}({\bf{x}}_{t},t) as the one previously used by Ho et al. , Song et al. . We use pretrained models from Ho et al. for CIFAR10, LSUN Bedrooms and Outdoor Churches, and from Song et al. for CelebA.

General setting

We follow the linear selection procedure to select a subsequence of timesteps τS⊂T\tau_{S}\subset T for all the datasets except CIFAR10, i.e., we select timesteps such that τi=⌊ci⌋\tau_{i}=\lfloor ci\rfloor for some cc. For CIFAR10, we select timesteps such that τi=⌊ci2⌋\tau_{i}=\lfloor ci^{2}\rfloor for some cc. The constant cc is selected so that τ−1\tau_{-1} is close to TT. We use Anderson acceleration as our fixed-point solver for all the experiments. We set the exiting equilibrium error of solver to 1e-3 and set the history length to 5. We allow a maximum of 15 solver forward steps in all the experiments with DEQ-DDIM, and use a maximum of 50 solver forward steps for DEQ-sDDIM. Finally, we use PyTorch’s inbuilt DataParallel module to handle parallelization. We use upto 4 NVIDIA Quadro RTX 8000 or RTX A6000 GPUs for all our experiments.

Training details for model inversion

We implement and test the code in PyTorch version 1.11.0. We use the Adam optimizer with a learning rate of 0.01. We train DEQs for 400 epochs on CIFAR10 and CelebA, and for 500 epochs on LSUN Bedroom, and LSUN Outdoor Church. The baseline is trained for 1000 epochs on CIFAR10 with T=100T=100, for 3000 epochs on CIFAR10 with T=10T=10, for 2500 epochs on CelebA, and for 2000 epochs on LSUN Bedroom, and LSUN Outdoor Church. At the beginning of inversion procedure with DEQ-DDIM, we sample xT∼N(0,I){\bf{x}}_{T}\sim{\cal{N}}(\mathbf{0},{\bf{I}}), and initialize the latents at all (or the subsequence of) timesteps to this value. For inversion with DEQ-sDDIM we also sample ϵ1:T∼N(0,I){\bm{\epsilon}}_{1:T}\sim{\cal{N}}(\mathbf{0},{\bf{I}}). We stop the training as soon as the loss falls below 0.5 for CIFAR10, and below 2 for other datasets.

Evaluation

We compute FID scores using the code provided by https://github.com/w86763777/pytorch-gan-metrics. We also use the precomputed statistics for CIFAR10, LSUN Bedrooms and Outdoor Churches provided in this github repository. For CelebA, we compute our own dataset statistics, as the precomputed statistics for images of resolution 64×\times64 are not included in this repository. While computing the statistics, we preprocess the images of CelebA in exactly the same way as done by Song et al. .

Appendix B Pseudocode

We provide PyTorch-style pseudocode to invert DDIM with DEQ approach in Algorithm box 3. Note that we use phantom gradient to compute inexact gradients.

Appendix C Ablation Studies for Model Inversion with DEQ-DDIM

We study the effect of the length of the diffusion chain on the convergence rate of optimization for model inversion in Fig. 6. We note that for sequential sampling, loss decreases slightly faster for the smaller diffusion chain (τS=10\tau_{S}=10 and 2020) than the longer one (τS=100\tau_{S}=100) for the baseline. However, for DEQ-DDIM, the length of diffusion chain doesn’t seem to have an effect on the rate of convergence as the loss curves for τS=100\tau_{S}=100 and τS=10\tau_{S}=10 and 2020 are nearly identical.

All the images in Fig. 4 are sampled with a subsequence of timesteps τS⊂T\tau_{S}\subset T, i.e., the number of latents in the diffusion chain used for training and the number of timesteps used for sampling an image from the recovered x^T\hat{x}_{T} were equal. We investigate if sampling with τS=T=1000\tau_{S}=T=1000 results in images with a better perceptual quality for the baseline. We display the recovered images for LSUN Bedrooms in Fig. 7. The length of diffusion chain during training time is τS=10\tau_{S}=10. We note that using more sampling steps does not result in inverted images that are closer to the original image. In some cases, samples generated with more sampling steps have some additional artifacts that are not present in the original image.

Comparing exact vs inexact gradients for backward pass of DEQ-DDIM

The choice of gradient calculation for the backward pass of DEQ affects both the training stability and convergence of DEQ-DDIM. Here, we compare the performance of the exact gradients and inexact gradients. Computing the inverse of Jacobian in Eq. 15 is difficult because the Jacobian can be prohibitively large. We follow Bai et al. to compute exact gradients using the following linear system

We use Broyden’s method to efficiently solve for v⊤{\bm{v}}^{\top} in this linear system. We compare it againt inexact gradients i.e., Jacobian free gradient used in Algorithm 2. We observe that training DEQ-DDIM with exact gradients becomes increasingly unstable as the training proceeds, especially for larger learning rates like 0.005. However, we can converge faster with larger learning rates like 0.01 with inexact gradients.

Effects of choice of initialization on convergence of DEQ-DDIM

The choice of initialization is critical for fast convergence of DEQs. Bai et al. initialize the initial estimate of fixed point of DEQ with zeros. However, in this work, we initialize all the latent states with xT{\bf{x}}_{T}, as it results in faster convergence shown in Figure 9. While both the initialization schemes converge to high quality images eventually, we observe that initializing with xT{\bf{x}}_{T} results in up to 3×3\times faster convergence compared to zero initialization. We observe a significant qualitative difference in the visualization of the intermediate states of the diffusion chain at different solver steps for the two initialization schemes as observed in Figure 10.

Appendix D Extending DEQ formulation to stochastic DDIM (DEQ-sDDIM)

A more general and highly stochastic generative process is to sample and integrate noises every time step, given by :

where ϵt∼N(0,I)\epsilon_{t}\sim{\cal{N}}(\textbf{0},{\bf{I}}). Here, σt=0\sigma_{t}=0 corresponds to a deterministic DDIM while σt=1−αt−11−αt1−αtαt−1\sigma_{t}=\sqrt{\dfrac{1-\alpha_{t-1}}{1-\alpha_{t}}}\sqrt{1-\dfrac{\alpha_{t}}{\alpha_{t-1}}} corresponds to a stochastic DDIM. Empirically, this is parameterized as σt(η)=η1−αt−11−αt1−αtαt−1\sigma_{t}(\eta)=\eta\sqrt{\dfrac{1-\alpha_{t-1}}{1-\alpha_{t}}}\sqrt{1-\dfrac{\alpha_{t}}{\alpha_{t-1}}} where η\eta is a hyperparameter to control stochasticity. Note that η=1\eta=1 corresponds to a DDPM .

Rearranging the terms in Eq. (20), we get

Let c1(t)=1−αt−1−σt2−αt−1(1−αt)αtc_{1}^{(t)}=\sqrt{1-\alpha_{t-1}-\sigma_{t}^{2}}-\sqrt{\dfrac{\alpha_{t-1}(1-\alpha_{t})}{\alpha_{t}}}. Then we can write

By induction, we can rewrite the above equation as:

This again defines a “fully-lower-triangular” inference process, where the update of xt{\bf{x}}_{t} depends on the noise prediction network ϵθ\epsilon_{\theta} applied to all subsequent states xt+1:T{\bf{x}}_{t+1:T}.

The above system of equations represent a DEQ with xT{\bf{x}}_{T} and ϵ1:T{\bm{\epsilon}}_{1:T} as input injection. We refer to this formulation of DEQ for stochastic DDIM with η>0\eta>0 as DEQ-sDDIM.

A major difference between the both is that DEQ-sDDIM can exploit the noises ϵ1:T{\bm{\epsilon}}_{1:T} sampled prior to fixed point solving as addition input injections. The insight here is that the noises along the sampling chain are independent of each other, thus allowing us to sample all the noises simultaneously and convert a highly stochastic autoregressive sampling process into a deterministic “fully-lower-triangular” DEQ.

where RootSolver(⋅\cdot) is any black-box fixed point solver.

D.2 Sample quality of images generated with DEQ-sDDIM

We verify that DEQ-sDDIM can generate images that are on par with the original DDIM by computing FID scores on 50,000 sampled images. We report our results in Table 2 and Table 4. We observe that our FID scores are comparable or slightly better to those from sequential DDIM. We also visualize images generated from the same latent xT{\bf{x}}_{T} at different levels of stochasticity controlled through η\eta. We display our generated images in Figure 12 and Figure 13.

D.3 Model inversion with DEQ-sDDIM

We report the minimum values of squared Frobenius norm between the recovered and target images averaged from 25 different runs in Figure 17. We use the same hyperparameters as the ones used for training DEQ models for DDIM in these experiments. DEQ-sDDIM is able to achieve low values of the reconstruction loss even for large values of η\eta like 11 as noted in Figure 17. We also plot training loss curves for different values of η\eta in Figure 17 on CIFAR10. We note that it indeed takes longer time to invert DEQ-sDDIM for higher values of η\eta. This is primarily because the fixed point solver needs more iterations to converge. However, despite that we obtain impressive model inversion results on CIFAR10 and CelebA. We visualize images generated with the recovered latent states in Figure 14 and Figure 15.

Appendix E Additional Qualitative Results

E.2 Model inversion with DEQ-DDIM