Ambient Diffusion: Learning Clean Distributions from Corrupted Data

Giannis Daras, Kulin Shah, Yuval Dagan, Aravind Gollakota, Alexandros G. Dimakis, Adam Klivans

Introduction

Diffusion generative models are emerging as versatile and powerful frameworks for learning high-dimensional distributions and solving inverse problems . Numerous recent developments have led to text conditional foundation models like Dalle-2 , Latent Diffusion and Imagen with incredible performance in general image domains. Training these models requires access to high-quality datasets which may be expensive or impossible to obtain. For example, direct images of black holes cannot be observed and high-quality MRI images require long scanning times, causing patient discomfort and motion artifacts .

Recently, showed that diffusion models can memorize examples from their training set. Further, an adversary can extract dataset samples given only query access to the model, leading to privacy, security and copyright concerns. For many applications, we may want to learn the distribution but not individual training images e.g. we might want to learn the distribution of X-ray scans but not memorize images of specific patient scans from the dataset. Hence, we may want to introduce corruption as a design choice. We show that it is possible to train diffusions that learn a distribution of clean data by only observing highly corrupted samples.

Prior work in supervised learning from corrupted data. The traditional approach to solving such problems involves training a restoration model using supervised learning to predict the clean image based on the measurements . The seminal Noise2Noise work introduced a practical algorithm for learning how to denoise in the absence of any non-noisy images. This framework and its generalizations have found applications in electron microscopy , tomographic image reconstruction , fluorescence image reconstruction , blind inverse problems , monocular depth estimation and proteomics . Another related line of work uses Stein’s Unbiased Risk Estimate (SURE) to optimize an unbiased estimator of the denoising objective without access to non-noisy data . We stress that the aforementioned research works study the problem of restoration, whereas are interested in the problem of sampling from the clean distribution. Restoration algorithms based on supervised learning are only effective when the corruption level is relatively low . However, it might be either not possible or not desirable to reconstruct individual samples. Instead, the desired goal may be to learn to generate fresh and completely unseen samples from the distribution of the uncorrupted data but without reconstructing individual training samples.

Indeed, for certain corruption processes, it is theoretically possible to perfectly learn a distribution only from highly corrupted samples (such as just random one-dimensional projections), even though individual sample denoising is usually impossible in such settings. Specifically, AmbientGAN showed that general dd dimensional distributions can be learned from scalar observations, by observing only projections on one-dimensional random Gaussian vectors, in the infinite training data limit. The theory requires an infinitely powerful discriminator and hence does not apply to diffusion models.

Our contributions. We present the first diffusion-based framework to learn an unknown distribution D{\cal D} when the training set only contains highly-corrupted examples drawn from D{\cal D}. Specifically, we consider the problem of learning to sample from the target distribution p0(x0)p_{0}({\bm{x}}_{0}) given corrupted samples Ax0A{\bm{x}}_{0} where A∼p(A)A\sim p(A) is a random corruption matrix (with known realizations and prior distribution) and x0∼p0(x0){\bm{x}}_{0}\sim p_{0}({\bm{x}}_{0}). Our main idea is to introduce additional measurement distortion during the diffusion process and require the model to predict the original corrupted image from the further corrupted image.

We use our algorithm to train diffusion models on standard benchmarks (CelebA, CIFAR-10 and AFHQ) with training data at different levels of corruption.

Given the learned conditional expectations we provide an approximate sampler for the target distribution p0(x0)p_{0}({\bm{x}}_{0}).

We show that for up to 90%90\% missing pixels, we can learn reasonably well the distribution of uncorrupted images. We outperform the previous state-of-the-art AmbientGAN and natural baselines.

We show that our models perform on par or even outperform state-of-the-art diffusion models for solving certain inverse problems even without ever seeing a clean image during training. Our models do so with a single prediction step while our baselines require hundreds of diffusion steps.

We use our algorithm to finetune foundational pretrained diffusion models. Our finetuning can be done in a few hours on a single GPU and we can use it to learn distributions with a few corrupted samples.

We show that models trained on sufficiently corrupted data do not memorize their training set. We measure the tradeoff between the amount of corruption (that controls the degree of memorization), the amount of training data and the quality of the learned generator.

We open-source our code and models: https://github.com/giannisdaras/ambient-diffusion.

Background

showed that we can learn the score function at level tt by optimizing for the score-matching objective:

Specifically, the score function can be written in terms of the minimizer of this objective as:

Inspired by this restoration interpretation of diffusion models, the Soft/Cold Diffusion works generalized diffusion models to look at non-Markovian corruption processes: xt=Ctx0+σtη{\bm{x}}_{t}=C_{t}{\bm{x}}_{0}+\sigma_{t}\bm{\eta}. Specifically, Soft Diffusion proposes the Soft Score Matching objective:

and shows that it is sufficient to recover the score function via a generalized Tweedie’s Formula:

For these generalized models, the matrix CtC_{t} is a design choice (similar to how we could choose the functions f,g{\bm{f}},g). Most importantly, for t=0t=0, the matrix CtC_{t} becomes the identity matrix and the noise σt\sigma_{t} becomes zero, i.e. we observe samples from the true distribution.

Method

For the sake of clarity, we first consider the case of random inpainting. If the image x0{\bm{x}}_{0} is viewed as a vector, we can think of the matrix AA as a diagonal matrix with ones in the entries that correspond to the preserved pixels and zeros in the erased pixels. We assume that p(A)p(A) samples a matrix where each entry in the diagonal is sampled i.i.d. with a probability 1−p1-p to be 11 and pp to be zero.

We would like to train a function hθ{\bm{h}}_{\theta} which receives a corruption matrix AA and a noisy version of a corrupted image, yt=A(x0+σtη)⏟xt{\bm{y}}_{t}=A\underbrace{({\bm{x}}_{0}+\sigma_{t}\bm{\eta})}_{{\bm{x}}_{t}} where η∼N(0,I)\bm{\eta}\sim\mathcal{N}(0,I), and produces an estimate for the conditional expectation. The simplest idea would be to simply ignore the missing pixels and optimize for:

Instead, we propose to further corrupt the samples before feeding them to the model, and ask the model to predict the original corrupted sample from the further corrupted image.

The key idea behind our algorithm is as follows: the learner does not know if a missing pixel is missing because we never had it (and hence do not know the ground truth) or because it was deliberately erased as part of the further corruption (in which case we do know the ground truth). Thus, the best learner cannot be inaccurate in the unobserved pixels because with non-zero probability it might be evaluated on some of them. Notice that the trained model behaves as a denoiser in the observed pixels and as an inpainter in the missing pixels. We also want to emphasize that the probability δ\delta of further corruption can be arbitrarily small as long as it stays positive.

2 Sampling

This idea works surprisingly well. Unless mentioned otherwise, we use it for all the experiments in the main paper and we show that we can generate samples that are reasonably close to the true distribution (as shown by metrics such as FID and Inception) even with 90%90\% of the pixels missing.

Sampling with Reconstruction Guidance. In the Fixed Mask Sampler, at any time tt, the prediction is a convex combination of the current value and the predicted denoised image. As t→0t\to 0, γt→0\gamma_{t}\to 0. Hence, for the masked pixels, the fixed mask sampler outputs the conditional expectation of their value given the observed pixels. This leads to averaging effects as the corruption gets higher. To correct this problem, we add one more term in the update: the Reconstruction Guidance term. The issue with the previous sampler is that the model never sees certain pixels. We would like to evaluate the model using different masks. However, the model outputs for the denoised image might be very different when evaluated with different masks. To account for this problem, we add an additional term that enforces updates that lead to consistency on the reconstructed image. The update of the sampler with Reconstruction Guidance becomes:

This sampler is inspired by the Reconstruction Guidance term used in Imagen to enforce consistency and correct for the sampling drift caused by imperfect score matching . We see modest improvements over the Fixed Mask Sampler for certain corruption ranges. We ablate this sampler in the Appendix, Section E.3.

Theory

Two simple examples that fit into this framework (see Corollaries A.1 and A.2 in the Appendix) are:

Experimental Evaluation

We first evaluate the restoration performance of our model for the task it was trained on (random inpainting and noise). We compare with state-of-the-art diffusion models that were trained on clean data. Specifically, for AFHQ we compare with the state-of-the-art EDM model and for CelebA we compare with DDIM . These models were not trained to denoise, but we can use the prior learned in the denoiser as in to solve any inverse problem. We experiment with the state-of-the-art reconstruction algorithms: DDRM and DPS .

We summarize the results in Table 1. Our model performs similarly to other diffusion models, even though it has never been trained on clean data. Further, it does so by requiring only one step, while all the baseline diffusion models require hundreds of steps to solve the same task with inferior or comparable performance. The performance of DDRM improves with more function evaluations at the cost of more computation. For DPS, we did not observe significant improvement by increasing the number of steps to more than 100100. We include results with noisy inpainted measurements and comparisons with a supervised method in the Appendix, Section E, Tables 3, 4. We want to emphasize that all the baselines we compare against have an advantage: they are trained on uncorrupted data. Instead, our models were only trained on corrupted data. This experiment indicates that: i) our training algorithm for learning the conditional expectation worked and ii) that the choice of corruption that diffusion models are trained to reverse matters for solving inverse problems.

Next, we evaluate the performance of our diffusion models as generative models. To the best of our knowledge, the only generative baseline with quantitative results for training on corrupted data is AmbientGAN which is trained on CIFAR-10. We further compare with a diffusion model trained without our further corruption algorithm. We plot the results in Figure 4. The diffusion model trained without our further corruption algorithm performs well for low corruption levels but collapses entirely for high corruption. Instead, our model trained with further corruption maintains reasonable corruption scores even for high corruption levels, outperforming the previous state-of-the-art AmbientGAN for all ranges of corruption levels.

For CelebA-HQ and AFHQ we could not find any generative baselines trained on corrupted data to compare against. Nevertheless, we report FID and Inception Scores and summarize our results in Table 4 to encourage further research in this area. As shown in the Table, for CelebA-HQ and AFHQ, we manage to maintain a decent FID score even with 90%90\% of the pixels deleted. For CIFAR-10, the performance degrades faster, potentially because of the lower resolution of the training images.

2 Finetuning foundation models on corrupted data

We can apply our technique to finetune a foundational diffusion model. For all our experiments, we use Deepfloyd’s IF model , which is one of the most powerful open-source diffusion generative models available. We choose this model over Stable Diffusion because it works in the pixel space (and hence our algorithm directly applies).

We show that we can finetune a foundational model on a limited dataset without memorizing the training examples. This experiment is motivated by the recent works of that show that diffusion generative models memorize training samples and they do it significantly more than previous generative models, such as GANs, especially when the training dataset is small. Specifically, train diffusion models on subsets of size {300,3000,30000}\{300,3000,30000\} of CelebA and they show that models trained on 300300 or 30003000 memorize and blatantly copy images from their training set.

We replicate this training experiment by finetuning the IF model on a subset of CelebA with 30003000 training examples. Results are shown in Figure 1. Standard finetuning of Deepfloyd’s IF on 30003000 images memorizes samples and produces almost exact copies of the training set. Instead, if we corrupt the images by deleting 80%80\% of the pixels prior to training and finetune, the memorization decreases sharply and there are distinct differences between the generated images and their nearest neighbors from the dataset. This is in spite of finetuning until convergence.

To quantify the memorization, we follow the methodology of . Specifically, we generate 10000 images from each model and we use DINO -v2 to compute top-11 similarity to the training images. Results are shown in Figure 6. Similarity values above 0.950.95 roughly correspond to the same person while similarities below 0.750.75 typically correspond to random faces. The standard finetuning (Red) often generates images that are near-identical with the training set. Instead, fine-tuning with corrupted samples (blue) shows a clear shift to the left. Visually we never observed a near-copy generated from our process – see also Figure 1.

We repeat this experiment for models trained on the full CelebA dataset and at different levels of corruption. We include the results in Figure 8 of the Appendix. As shown, the more we increase the corruption level the more the distribution of similarities shifts to the left, indicating less memorization. However, this comes at the cost of decreased performance, as reported in Table 4.

New domains and different corruption. We show that we can also finetune a pre-trained foundation model on a new domain given a limited-sized dataset in a few hours in a single GPU. Figure 5 shows generated samples from a finetuned model on a dataset containing 155155 examples of brain tumor MRI images . As shown, the model learns the statistics of full brain tumor MRI images while only trained on brain-tumor images that have a random box obfuscating 25%25\% of the image. The training set was resized to 64×6464\times 64 but the generated images are at 256×256256\times 256 by simply leveraging the power of the cascaded Deepfloyd IF.

Acknowledgements.

The authors would like to thank Tom Goldstein for insightful discussions that benefited this work. This research has been supported by NSF Grants CCF 1763702, AF 1901292, CNS 2148141, Tripods CCF 1934932, NSF AI Institute for Foundations of Machine Learning (IFML) 2019844, the Texas Advanced Computing Center (TACC) and research gifts by Western Digital, WNCG IAP, UT Austin Machine Learning Lab (MLL), Cisco and the Archie Straiton Endowed Faculty Fellowship. Giannis Daras has been supported by the Onassis Fellowship (Scholarship ID: F ZS 012-1/2022-2023), the Bodossaki Fellowship and the Leventis Fellowship.

References

Appendix A Proofs

Let hθ∗{\bm{h}}_{\theta^{*}} be a minimizer of equation 3.2, and for brevity let

be the difference between hθ∗{\bm{h}}_{\theta^{*}} and the claimed optimal solution. We will now argue that f{\bm{f}} must be identically zero.

Here the first term is the irreducible error, while the third term vanishes by the tower law of expectations:

This intuition can be formalized as follows. Fix a distribution pA(A)p_{A}(A) over corruption matrices. For a distribution p0(x0)p_{0}({\bm{x}}_{0}), denote by corrupt⁡(pA,p0)\operatorname{corrupt}(p_{A},p_{0}) the distribution over pairs (A,Ax0)(A,A{\bm{x}}_{0}) where A∼pA(A)A\sim p_{A}(A) and x0∼p0(x0){\bm{x}}_{0}\sim p_{0}({\bm{x}}_{0}). We say that it is possible to reconstruct p0(x0)p_{0}({\bm{x}}_{0}) from random corruptions A∼pA(A)A\sim p_{A}(A) if the following holds: for any two distributions, p0(x0)p_{0}({\bm{x}}_{0}) and p0′(x0′)p_{0}^{\prime}({\bm{x}}_{0}^{\prime}) that satisfy Assumptions 1-3 of , if corrupt⁡(pA,p0)=corrupt⁡(pA,p0′)\operatorname{corrupt}(p_{A},p_{0})=\operatorname{corrupt}(p_{A},p_{0}^{\prime}), then p0=p0′p_{0}=p_{0}^{\prime}. Similarly, we say that it is possible to reconstruct p0(x0)p_{0}(x_{0}) from conditional expectations given A∼p(A)A\sim p(A) if the following holds: for any distribution p0(x0)p_{0}({\bm{x}}_{0}) and p0′(x0′)p_{0}^{\prime}({\bm{x}}_{0}^{\prime}) that satisfy Assumptions 1-3 of , if for all xx, tt and AA in the support of pAp_{A},

then p0=p0′p_{0}=p_{0}^{\prime}. Here, p0,t(x0,xt)p_{0,t}({\bm{x}}_{0},{\bm{x}}_{t}) is obtained by sampling x0∼p0{\bm{x}}_{0}\sim p_{0} and xt=x0+σtη{\bm{x}}_{t}={\bm{x}}_{0}+\sigma_{t}\bm{\eta} where η∼N(0,I)\bm{\eta}\sim\mathcal{N}(0,I). Similarly, p0,t′(x0′,xt′)p^{\prime}_{0,t}({\bm{x}}_{0}^{\prime},{\bm{x}}_{t}^{\prime}) is obtained by the same process where x0′{\bm{x}}_{0}^{\prime} is instead sampled from p0′p_{0}^{\prime}. We state the following lemma:

Appendix B Broader Impact and Risks

Generative models in general hold the potential to have far-reaching impacts on society in a variety of forms, coupled with several associated risks . Among other potential applications, they can be utilized to create deceptive images and perpetuate societal biases. To the best of our knowledge, our paper does not amplify any of these existing risks. Regarding the included MRI results, we want to clarify that we make no claim that such results are diagnostically useful. This experiment serves only as a toy demonstration that it can be potentially feasible to learn the distribution of MRI scans with corrupted samples. Significant further research must be done in collaboration with radiologists before our algorithm gets tested in clinical trials. Finally, we want to underline again that even though our approach seems to mitigate the memorization issue in generative models, we cannot guarantee the privacy of any training sample unless we make assumptions about the data distribution. Hence, we strongly discourage using this algorithm in applications where privacy is important before this research topic is investigated further.

Appendix C Training Details

We open-source our code and models to facilitate further research in this area: https://github.com/giannisdaras/ambient-diffusion.

We trained models from scratch at different corruption levels on CelebA-HQ, AFHQ and CIFAR-10. The resolution of the first two datasets was set to 64×6464\times 64 and for CIFAR-10 we trained on 32×3232\times 32.

We started from EDM’s official implementation and made some necessary changes. Architecturally, the only change we made was to replace the convolutional layers with Gated Convolutions that are known to perform well for inpainting problems. We observed that this change stabilized the training significantly, especially in the high-corruptions regime. As in EDM, we use the architecture from the DDPM++ paper.

To avoid additional design complexity, we tried to keep our hyperparameters as close as possible to the EDM paper. We observed that for high corruption levels, it was useful to add gradient clipping, otherwise, the training would often diverge. For all our experiments, we use gradient clipping with max-norm set to 1.01.0. We underline that unfortunately, even with gradient clipping, the training at high corruption levels (p≥0.8p\geq 0.8), still diverges sometimes. Whenever this happened, we restarted the training from an earlier checkpoint. We list the rest of the hyperparameters we used in Table 2.

Training diffusion models from scratch is quite computationally intensive. We trained all our models for 200000200000 iterations. Our CIFAR-10 models required ≈2\approx 2 days of training each on 66 A100 GPUs. Our AFHQ and CelebA-HQ models required ≈6\approx 6 days of training each on 66 A100 GPUs. These numbers roughly match the performance reported in the EDM paper, indicating that the extra corruption we need to do on matrix AA does not increase training time.

Due to the increased computational complexity of training these models, we could not extensively optimize the hyperparameters, e.g. the δ\delta probability in our extra corruption. For higher corruption, e.g. for p=0.9p=0.9, we noticed that we had to increase δ\delta in order for the model to learn to perform well on the unobserved pixels.

Finetuning Deepfloyd’s IF.

We access Deepfloyd’s IF model through the diffusers library. The model is a Cascaded Diffusion Model . The first part of the pipeline is a text-conditional diffusion model that outputs images at resolution 64×6464\times 64. Next in the pipeline, there are two diffusion models that are conditioned both in the input text and the low-resolution output of the previous stage. The first upscaling module increases the resolution from 64×6464\times 64 to 256×256256\times 256 and the final one from 256×256256\times 256 to 1024×10241024\times 1024.

To reduce the computational requirements of the finetuning, we only finetune the first text-conditional diffusion model that works with 64×6464\times 64 resolution images. Once the finetuning is completed, we use again the whole model to generate high-resolution images.

For our CelebA finetuning experiments, we set δ=0.1\delta=0.1 and p=0.8p=0.8. We experiment with the full training set, a subset of size 30003000 (see Figure 1) and a subset of 300300. For the model trained with only 300300 heavily corrupted samples, we did not observe memorization but the samples were of very low quality. Intuitively, our algorithm provides a way to control the trade-off between memorization and fidelity. Fully exploring this trade-off is a very promising direction for future work. For our MRI experiments, we use two non-overlapping blocks that each obfuscate 25%25\% of the image and we evaluate the model in one of them.

All of our fine-tuning experiments can be completed in a few hours. Training for 1500015000 iterations takes ≈10\approx 10 hours on an A100 GPU, but we usually get the best checkpoints earlier in the training.

Appendix D Evaluation Details

Our FID score is computed with respect to the training set, as is standard practice, e.g. see . For each of our models trained from scratch, we generate 5000050000 images using the seeds 0−499990-49999. Once we generate our images, we use the code provided in the official implementation of the EDM paper for the FID computation.

Appendix E Additional Experiments

In Table 1 of the main paper, we compare the restoration performance of our models and vanilla diffusion models (trained with uncorrupted images). We compare the restoration performance in the task of random inpainting because it is straightforward to use our models to solve this inverse problem. It is potentially feasible to use our trained generative models to solve any (linear or non-linear) inverse problem but we leave this direction for future work.

We use the EDM state-of-the-art model trained on AFHQ as our baseline (as we did in the main paper). To use this pre-trained model to solve the noisy random inpainting inverse problem, we need a reconstruction algorithm. We experiment again with DPS and DDRM which can both handle inverse problems with noise in the measurements. We present our results in Table 3. As shown, our models significantly outperform the EDM models that use the DDRM reconstruction algorithm and perform on par with the EDM models that use the DPS reconstruction algorithm.

E.2 Comparison with Supervised Methods

For completeness, we include a comparison with Masked AutoEncoders , a state-of-the-art supervised method for solving the random inpainting problem. The official repository of this paper does not include models trained on AFHQ. We compare with the available models that are trained on the iNaturalist dataset which is the most semantically close dataset we could find. We emphasize that this model was trained with access to uncorrupted images. Results are shown in 4. As shown, our method and DPS outperform this supervised baseline. We underline that this experiment is included for completeness and does not exclude the possibility that there are more performant supervised alternatives for random inpainting.

E.3 Sampler Ablation

By Lemma A.3, if it is possible to learn p0(x0)p_{0}({\bm{x}}_{0}) using corrupted samples then it is also possible to use our learned model to sample from p0(x0)p_{0}({\bm{x}}_{0}). Even though such an algorithm exists, we do not know which one it is.

In the paper, we proposed to simple ideas for sampling, the Fixed Mask Sampler and the Reconstruction Guidance Sampler. The Fixed Mask Sampler fixes a mask throughout the sampling process. Hence, sampling with this algorithm is equivalent to first sampling some pixels from the marginals of p0(x0)p_{0}({\bm{x}}_{0}) and then completing the rest of the pixels with the best reconstruction (the conditional expectation) in the last step. This simple sampler performs remarkably well and we use it throughout the main paper.

To demonstrate that we can further improve the sampling performance, we proposed the Reconstruction Guidance sampler that takes into account all the pixels by forcing predictions of the model with different masks to be consistent. In Table 5 we ablate the performance of this alternative sampler. For the reconstruction guidance sampler, we select each time four masks at random and we add an extra update term to the Fixed Mask Sampler that ensures that the predictions of the Fixed Mask Sampler are not very different compared to the predictions with the other four masks (that have different context regarding the current iterate xt{\bm{x}}_{t}). We set the guidance parameter wtw_{t} to the value 5e−45e-4. As shown, this sampler improves the performance, especially for the low corruption probabilities where the extra masks give significant information about the current state to the predictions given only one fixed mask. However, the benefits of this sampler are vanishing for higher corruption. Given that the two samples perform on par and that the Reconstruction Guidance Sampler is much more computationally intensive (we need one extra prediction per step for each extra mask), we choose to use the Fixed Mask Sampler for all the experiments in the paper.

E.4 Additional Figures

Figure 7 shows reconstructions of AFHQ corrupted images with the EDM AFHQ model trained on clean data (columns 3, 4) and our model trained on corrupted data (column 5). The restoration task is random inpainting at probability p=0.8p=0.8. The last two rows also have measurement noise with σy0=0.05\sigma_{{\bm{y}}_{0}}=0.05.

In Figure 8, we repeat the experiment of Figure 6 of the main paper but for models trained on the full CelebA dataset and at different levels of corruption. As shown, increasing the corruption level leads to a clear shift of the distribution to the left, indicating less memorization. This comes at the cost of decreased performance, as reported in Table 4.

In the remaining pages, we include uncurated unconditional generations of our models trained at different corruption levels pp. Results are shown in Figures 9,10,11.