Pruning then Reweighting: Towards Data-Efficient Training of Diffusion Models
Yize Li, Yihua Zhang, Sijia Liu, Xue Lin
I Introduction
Diffusion Models (DMs) belong to a recent class of generative models, which have achieved state-of-the-art generation performance . Despite their superiority in terms of training stability, versatility, and scalability, DMs are known for their slow generation speeds due to the requirement of reverse diffusion processing by passing through the generator at massive times. Consequently, there is considerable interest in enhancing the inference speed of DMs . Furthermore, DMs are recognized for their high training costs. Modeling complicated and high-dimensional data distributions requires numerous iterations, resulting in exponential growth in training costs under the increasing resolution and diversity of the data.
Several works have considered speeding up diffusion training by the progressive patch size , masked patches , momentum stochastic gradient descent (SGD) and a clamped signal-to-noise ratio (SNR) weight at time-step . However, none of them attempted to achieve efficient training through the lens of dataset pruning (or coreset selection). To the best of our knowledge, this is the first work to investigate how the coreset size of training data influences the generation ability of DMs. In this study, we first utilize a GAN-based data selection method for diffusion training, which consists of feature embedding and data scoring. To refine training data distribution, a perceptually aligned embedding function , such as the latent space of a pre-trained image classifier (e.g., Inceptionv3 ) is to acquire the data feature space. Then, a scoring criterion (e.g., Gaussian model) is to rank each data point in the embedding space and remove less relevant data. Nevertheless, we discover that such a data selection approach may generalize poorly to DMs on small-scale datasets. Hence, there is a pressing need for innovations to enhance the current data selection scheme for diffusion-based generative models.
We summarize our proposed pipeline in Fig. 1, which investigates the encoder and scoring method to implement data selection in DMs. Inceptionv3 , ResNet-18 , CLIP and DDAE are adopted as the choices of surrogate models (encoder) and the scoring functions (dataset pruning methods) are Gaussian model and Moderate-DS , which keep data points with scores within the scoring threshold. One key observation is that simply pruning the dataset might lower generation capability, with the generative capacity of each class decreasing to varying extents in Fig. 2. To address this issue, we leverage a class-wise reweighting strategy by distributionally robust optimization (DRO ), to optimize the class weights that are dynamically updated according to the marginal loss on each class. Experimental results on the pixel-level DDPM , the latent-level Masked Diffusion Transformer (MDT) and Stable Diffusion (SD) demonstrate that our method could accelerate diffusion training from 2.34 up to 8.32 while maintaining comparable or even superior generation ability. The main contributions are highlighted below.
We investigate the problem of efficient DM training through the lens of dataset pruning for the first time, which selects coreset from the latent space through surrogate models.
We develop a novel class-wise reweighting strategy to enhance generation capacity by minimizing the variance between the target proxy model and the reference model.
We achieve comparable performances on DDPM and notable sampling improvements on latent diffusion models (LDMs) while obtaining gains in computation efficiency.
II Related Work
Diffusion models are proposed to capture the high-dimensional nature of data distributions, which are dominating a new era by exceeding the Generative adversarial networks (GANs) . The backbone networks of DMs generally include the convolutional U-Net , and the transformer-based architectures with attention layers.
The sampling of DMs is typically costly because of the iterative denoising process with UNet and the DM training is always time-consuming by massive steps. To address these issues, existing works concentrate on reducing sampling steps through step distillation and efficient sampling solvers, including DDIM and DPM-Solver . Other recent works consider compression and utilize the property of the model architecture . Furthermore, accelerating diffusion training is achieved by gradually scaling up image size , or token merging and masking in transformer-based DMs.
II-B Dataset Pruning
Dataset pruning, also known as coreset selection, refers to reducing training data by creating a more compact dataset . A small representative subset can be approximated based on training dynamics as the score criterion , and loss or gradient perspectives, such as GRAD-MATCH , RHO-LOSS and InfoBatch .
III Methodology
Diffusion models. DMs include a forward noising process and a backward denoising process to estimate the distribution of data iteratively . Given the clean input , it is gradually turned into the noisy over time steps ( at each time step ) by Gaussian noise in the forward diffusion process. In the backward sampling process, a noisy sample is progressively denoised to generate an uncorrupted output. The objective of DM can be simplified by minimizing the noise approximation error
where represents the noise estimator at time step over trainable parameter regarding with the condition (e.g., class label or text prompt). In conditional DDPM , denotes the input image, while in latent diffusion model (LDM) , is the latent feature. Classifier-free guidance has been demonstrated to significantly enhance the sample quality of class-conditioned DMs. Specifically, a guidance weight is introduced to balance generation quality and sample diversity, where a conditional DM with the condition is jointly trained with an unconditional DM. The new noise estimation from Eq. (1) is formulated as .
III-B Dataset Pruning
III-C Class-wise Reweighting
Distinct differences in sampling abilities persist across all classes, as shown in Fig. 2. Class-wise reweighting aims to improve overall generative performance after dataset pruning by considering these differences between diverse domains. To acquire class weights, a proxy model is trained by the worst-case loss over classes, which follows a mini-max optimization as distributionally robust optimization (DRO):
IV Experiments
We vary the surrogate models as the encoder by Inceptionv3 , ResNet-18 , ResNet-50-based CLIP and DDAE (an unconditional DDPM) on three datasets with class labels, including CIFAR-10 with image size of 3232, ImageNet and ImageNette (a subset containing 10 easy classes from ImageNet) with image size of 256256. The pixel-wise DM is DDPM with classifier-free guidance and LDMs are MDT with the size of S, mask ratio 0.3 and Adan optimizer , and SD . DDPM is trained from scratch on 2000 epochs, MDT is trained on 60 epochs and SD is fine-tuned on 50 epochs. Generation quality is evaluated in Fréchet Inception Distance (FID) on 50k generated samples. Both DDPM and SD are efficiently inferenced by DDIM sampler with 100 and 50 steps, class classifier-free guidance as 0.3 and 5 respectively. MDT is evaluated with 250 DDPM sampling steps and 3.8 classifier-free guidance . Considering a more reliable and unbiased estimator of image quality on ImageNet, we adopt CMMD by Vision-Transformer-based CLIP embeddings and the maximum mean discrepancy distance with the Gaussian kernel. The DDPM proxy model on class reweighting follows the same setups mentioned above. The MDT proxy model follows similar settings except for 6 training epochs.
IV-B DDPM Results
We first investigate the data-efficient training of DMs on CIFAR-10 with low resolution (3232). In Table I, DMs trained on pruned dataset size ranging from 5000 (10% data ratio) to 20000 (40%) show different generative capabilities. We discover that the GAN-based instance selection is generalized poorly to DMs, where FID scores are even greatly worse than those of Uniform Random, which selects data for each class randomly in a unified way. The performance decline is from 4.90 up to 13.91 with the decrease in dataset size, proving that it is difficult to describe the compact feature space by such an encoder and data pruning method. Therefore, we execute tentative experiments on surrogate encoding models, including ResNet-18 , CLIP , and DDAE (an unconditional self-supervised learner based on DDPM) as image encoders. The scoring function is changed to Moderate-DS , choosing those data samples close to the median. As shown in Table I, ResNet-18 and DDAE yield more effective coresets compared to CLIP, which achieves lower FID than Uniform Random. The potential explanation, as supported by Table V, is that both ResNet-18 and DDAE are trained solely on CIFAR-10 without any data transformations aimed at enlarging image size. In contrast, CLIP is pre-trained on a larger dataset with higher-resolution images. Consequently, latent features from CLIP on CIFAR-10 are somewhat expanded and sub-optimal. Another observation is that different classes maintain diverse generative abilities, as depicted in Fig. 2. To tackle this problem, class-wise reweighting is leveraged to enhance the generation levels from a class-specific perspective. Equipped with class-wise reweighting, DDPM trained on the coreset selected by ResNet-18 and Moderate-DS achieves a 6.71 FID score under 10% of data and an FID score of 4.18 with 40% of training samples. Inspired by InfoBatch , we consider a pruned dataset training ratio of 0.875 in Annealing to improve further generation abilities, where DDPM is trained on the subset before the 87.5% training epoch and then on the full dataset until the end. We find that Annealing is significantly helpful in enhancing generative qualities especially when the dataset size is smaller and our approach outperforms InfoBatch in Table II. Due to the partial full data training, the training speed is affected by Annealing in Table III, showing the trade-offs between computational efficiency and image synthesis capacity.
IV-C MDT Results
Our dataset pruning is further extended to MDT on ImageNet. MDT learns the contextual relation among object semantic parts by masking certain tokens in the latent space. Gaussian model selects a more superior and compact subset than Moderate-DS via feature embedding from Inceptionv3. Furthermore, class-wise reweighting is general to all data pruning approaches and different pruning ratios. Remarkably, as shown in Table IV, MDT equipped with reweighting on merely 20% of data samples achieves a better FID score (13.94 vs. 17.11) than the model trained on the entire dataset. It demonstrates the possibility of both saving training costs and achieving qualified class-conditional image generation.
IV-D SD Results
We evaluate dataset pruning on ImageNette, a subset from ImageNet, by fine tuning SD (Stable Diffusion ‘v1-4’ ). The prompt for sampling is ‘a photo of a class name’. By computing FID on a total of 50k sampling images in 256256 resolution under classifier-free guidance, SD fine-tuned on only 40% of the data significantly surpasses the model on the entire dataset in Table V. An interesting finding is that all models fine-tuned on the subsets show even better generation capability, highlighting the potential data redundancy in large LDM fine-tuning. Note that class-wise reweighting is not applied in this case, because SD fine-tuned on the subset has surpassed the one on the entire dataset, and thus majority of margin loss from Eq. (2) is clipped to 0.
V Conclusion
In this work, we investigate data-efficient DM training by data selection and class reweighting. As the first study on data-pruned DM training, we demonstrate its remarkable robustness across various, reducing computing overhead by up to 8. Furthermore, we reveal the presence of training data redundancy in both pixel-level and latent-level DMs. Overall, we believe our findings and approach provide a solid foundation for building scalable and efficient artificial intelligence-generated content systems.