Catch-Up Distillation: You Only Need to Train Once for Accelerating Sampling
Shitong Shao, Xu Dai, Lujun Li, Huanran Chen, Yang Hu, Shouyi Yin
Introduction
Diffusion Probability Models (DPMs) , Variational Auto Encoders (VAEs) , and Generate Adversarial Networks (GANs) have achieved remarkable success across various applications, including image synthesis , audio synthesis , 3D reconstruction , and super-resolution . In recent years, DPMs, especially score-based probabilistic models , have emerged as the new state-of-the-art family of generative models. They demonstrate superior abilities to generate more coherent and diverse samples compared to their VAE and GAN counterparts . This attribute to DPMs’ theoretical completeness and exceptional image synthesis capabilities . have shown that DPMs at continuous time steps can be interpreted as a score function matching problem based on Stochastic Differential Equations (SDE) and Ordinary Differential Equations (ODE). Since the inception of this theory, research endeavors including SNIPS and Analytic-DPM have probed into SDE-based generative models, amplifying their applicability. The stochastic sampling employed by SDE-based models elevates the quality of synthetic samples over deterministic sampling. Nevertheless, these advanced designs also involve expensive computational budgets, posing significant challenges to resource-constrained research labs and practical applications in industry.
As a result, researchers have also explored new paradigms that incorporate SDE-based training and ODE-based sampling, which have effectively accelerated sampling as demonstrated in studies such as . Unfortunately, this paradigm is limited by the bottleneck of being training-free (i.e., can be directly served for acceleration in inference without training overhead), resulting in weaker accelerated sampling effects compared to training-dependent (i.e., need additional training overhead to achieve accelerating sampling) paradigms. A popular and widely-used training-dependent accelerated sampling paradigm involves the application of knowledge distillation to expedite the sampling process . This paradigm was first proposed in Progressive Distillation (PD) , with the core idea being to achieve accelerated sampling incrementally by distilling a multi-step process into a single step. After that, a range of works have been carried out to improve PD’s performance or extend its application scenarios. Among them, additionally distills the intermediate layer output of the noise estimation model, and extends the PD algorithm to conditional sampling scenarios. Just recently, Consistency Distillation was proposed to achieve the goal of the PD algorithm through a single additional training session, as opposed to multiple progressive training sessions.
However, all currently proposed distillation-based accelerated sampling algorithms have the following drawbacks: (a) They involve one or more additional training stages for distilling. (b) They need pre-training weights, and it must be ensured that the student’s architecture is identical to that of the teacher. (c) They only consider discrete time steps and thus cannot synthesize images taking advantage of arbitrary numerical integration algorithms (e.g., Euler–Maruyama method and Runge-Kutta method). These shortcomings have, to some extent, prevented the application of these distillation-based accelerated sampling algorithms to broader DPM paradigms. To address these issues, we propose Catch-Up Distillation (CUD), which treats the previous moment output of the velocity estimation model as the teacher output and the current moment output of the velocity estimation model as the student output, and applies Runge-Kutta-based multi-step alignment distillation (as illustrated in Fig. 1) to let the student output “catch up” with the teacher output while aligning the student output with the ground truth label. Based on this, CUD can complete accelerated sampling with a single training session, i.e., the original training session of DPMs, without requiring any pre-trained weights. Our contribution can be summarized as follows:
We present Catch-Up Distillation (CUD), the first accelerated sampling framework in a single training session without pre-training weights to extend the applicability of distillation-based accelerated sampling algorithms.
We search the design space of CUD and obtain some suitable strategies based on experiments and theories, which allow for a significant improvement in the quality of the synthetic samples.
We conduct extensive comparison and ablation experiments on the CIFAR-10, MNIST, and ImageNet-64 datasets to verify that CUD can achieve superior performance compared to the original DPM training, using the same number of iteration steps.
Background
where and MSE refer to the weight function that satisfies and Mean Square Error (MSE), respectively. is a very small amount (i.e. 1e-5) to empirically avoid unnecessary fitting overhead. Furthermore, is chosen to ensure that fits the same target velocity . After training the sampling can be done by definite integral, i.e., , s.t., . Intuitively, we can utilize this ODE solution approach to gradually reduce noise in a “clean” image, and restore it to the original data distribution, denoted as . In a series of well-established studies , numerical methods such as Euler’s method, Heun’s method, and Runge-Kutta method, have been effectively employed for solving ODEs in the sampling phase.
Reparameterized Noise Encoder.
Rectified flow employs a distillation technique to reduce transport costs by fitting instead of . Commonly, this method requires an additional training phase. Do not like vanilla Rectified flow, in , the authors propose defining a reparameterized noise encoder to reparameterize Gaussian noise, ensuring a smooth mapping from to and inducing efficient optimization. The new optimization objective can be expressed as
Here, and represent the loss weight (default as ) and Kullback-Leibler divergence, respectively. Although is primarily designed to minimize the curvature on the transport path, we demonstrate that it also reduces transport costs empirically. We substantiate this through Theorem 2.1, which implies that as the Mutual Information (MI) between two distributions increases, the corresponding transport cost diminishes. Higher MI signifies a stronger correlation between the two probability distributions. This means that as long as the cost function chosen for training is inversely proportional to MI, we can observe a smaller transport cost between samples drawn from these distributions than without a reparameterized noise encoder.
With the above analysis in mind, we will take the optimization objective 2 as a baseline and explore how to utilize distillation for accelerating sampling.
Knowledge Distillation in Accelerating Sampling.
Knowledge distillation is an effective technique for compressing models that can enhance the generalization ability of lightweight models. While distillation is commonly used to compress models, it can also accelerate the sampling process of DPMs. The studies of compress the multi-step sampling process of the DPM into a single step, effectively reducing computational overhead without compromising the quality of the synthetic images. However, these algorithms require pre-trained weights, and the student architecture is heavily dependent on the selected teacher architecture, limiting their scalability. In this paper, we aim to enable a DPM to function as both a teacher and a student, and accelerate sampling within a single training session without requiring any pre-trained weights, as compared to the above-mentioned work. It means that CUD can be considered a standard training session for DPMs with some additional loss terms, therefore ensuring its good portability.
Methodology
The conventional training objective for ODE/SDE-based generative models is to take input samples at different time points and output the noise , sample , or velocity , aligning them with the corresponding ground truth labels. This paradigm is incapable of allowing the model to generate high-quality samples in a few sampling steps scenario. Although Song et.al. propose Consistency Training (CT), the method is only applicable to empirical PF ODE since it needs Karra’s diffusion model paradigm for supervision (detailed explanation can be found in Appendix B) and has a huge performance gap compared to their proposed Consistency Distillation (CD). These shortcomings prevent CT from generalizing to other diffusion model paradigms, e.g., VP-SDE, VE-SDE. Considering this issue, our proposed Catch-Up Distillation (CUD) aims to adjust the original training paradigm so that it not only performs ground truth label alignment but also enables “catch up” with the output from catch-up sampling. The term “catch-up sampling” refers to using a numerical integral solver to estimate from , where “” refers to the step size of the discrete sampling. As presented in Fig. 1, CUD leverages Runge-Kutta-based multi-step alignment distillation for achieving accelerated sampling. Particularly, CUD also includes a series of simple but effective strategies (e.g. use the training model for catch-up sampling, random step size, dynamic skip connection) derived from searching the design space. Ultimately, the procedures of the CUD algorithm using Runge-Kutta 12, 23, and 34 are presented in Algorithms 3, C and C, respectively. Algorithms C and C and CUD’s limitations and broader impact can be found in the Appendix E and H.
Integrating CD into continuous time steps may seem beneficial, but as demonstrated in Table 1 in our experiments, it can lead to training collapse. This issue arises because the Exponential Moving Average (EMA) model, , which guides training, may become unreliable due to inaccurate updates. We can derive Theorem 3.1 and thus simply explicate this fact.
(Proof in Appendix C) Assume that and after training convergence, where , , and denote the norm, the parameters updated via EMA, and the pre-trained weight, respectively. And satisfies Lipschitz condition, i.e., , there exists such that for all . Then if exists, we can obtain that
The theorem establishes an upper bound, ensuring the stability of the distillation process. When the EMA model’s one-step update is inaccurate, will become larger, leading CD to training collapse. It also implies that the difference between and must be adequately small for the process to remain stable. In the standard diffusion model training session, we do not have a pre-trained weight like CD, so we have to replace with . In line with common distillation algorithms, such as vanilla KD , RKD , and HSAKD , the teacher model output and the ground truth label mutually supervise the student model output. This approach enhances the student model’s generalization ability empirically. Consequently, we introduce a ground truth label to enable effective supervision, thereby avoiding the EMA model’s misdirected guidance. The output of , for all , should align with . We term this minimal functional algorithm as basic CUD, which synchronizes the ground truth label while “catches up” to the model output at the previous moment. Thus, the new base loss can be denoted asWe omit the optimization of in all subsequent equations but utilize it in our experiments.
where and is calculated by catch-up sampling using ODE . And represent the loss weights, which can even be parameterized by the number of current iterations, i.e., dynamic weight in Table 1.
2 Runge-Kutta-Based Multi-Step Alignment Distillation
One of the key components of our CUD is catch-up sampling, which can be solved by different ODE solvers. Notably, Euler’s method and Heun’s method are the subsets of Runge-Kutta methods, w.r.t., Runge-Kutta 12, and Runge-Kutta 23. So in this study, we apply Runge-Kutta methods to model catch-up sampling on a generic perspective. We only consider the Runge-Kutta algorithm up to order 3, as higher-order algorithms would lead to excessive computational overhead, even though it is possible to ignore the backpropagation of gradients and parameter updates by means of inference’s form, e.g., torch.no_grad(). In general, the points sampled
by the higher-order Runge-Kutta will only serve for one sampling step. But this format obviously wastes a very large amount of training overhead, so is it possible to make the best use of all the sampling points? The answer is yes, as shown in Fig. 1. We can achieve Runge-Kutta-based multi-step alignment distillation, which aims to give the next step, the next-next step and even the next-next-next step of the estimated velocities simultaneously through all available sampling points, and then align them with outputs of different heads , where refers to the order of the Runge-Kutta method. Note that on here is equivalent to as mentioned earlier. As illustrated in Fig. 2, the input of all heads is the output of , and all heads are single meta-encoders that can be modeled as the sequence of consecutive GroupNorm-SiLU-Conv. By derivation in Appendix D, we can derive Runge-Kutta (23/34)-based multi-step alignment distillation as follows:
Compared to a simple one-step alignment distillation, the use of multi-step alignment distillation provides more comprehensive information about , resulting in preventing asynchronous model updates and improving model performance.
3 Investigate Design Space
To some extent, the basic CUD has facilitated accelerated sampling; however, further enhancements are feasible. In this subsection, we delineate the design space and conduct a comprehensive analysis, ultimately proposing strategies superior to the basic CUD.
The update of the EMA model can enhance the model’s generalization ability. Nevertheless, during the training phase, the disparity between the parameters of the EMA model and the training model may cause the upper bound of Theorem 3.1 to be excessively loose, particularly in relation to the term (as concluded in Appendix C). A more sensible and effective approach is to discard the EMA model and utilize the training model directly for catch-up sampling. In this case, the errors introduced by the incorrect estimation of the EMA model will no longer exist, further improving the performance of the model.
Random Step Size in Catch-Up Sampling.
For accelerating sampling, a large interval between sample points is employed to reduce the sample count. However, in widely-used diffusion models such as DDPM , NCSN , and Rectified flow, the interval between sampling points approaches zero, as these models necessitate accurate estimation of differential equations to prevent training collapse. Intuitively, utilizing a fixed value of for CUD may result in the model converging to suboptimal solutions. This occurs because, for , the model can only obtain the solution in the time point and is unable to explore solutions in other time points that could provide a more precise estimate of the ODE. To tackle this problem, we implement a non-fixed-step size strategy, uniform, implying that the new step size follows a uniform distribution ( is a predetermined step size, defaulting to 1/16). Additionally, we introduce a strategy named rule, which determines the new step size as , aiming to ensure the quality of the synthetic image retains a degree of ambiguity when , necessitating a larger step size to augment the distillation strength. Conversely, when
, the synthetic image’s quality is superior, thus requiring a smaller step size to maintain stability.
Dynamic Skip Connection.
Training a velocity estimation model with a fixed architecture directly would be far from ideal. As demonstrated in Appendix F, the cost of fitting the training model varies at different . In past work , they both apply a special operation that skips the entire model with residual mapping, i.e., , where and are two functions that input and output a scalar. However, this skip connection is coarse-grained and does not take advantage of the skip connection that contemporary mainstream neural networks have themselves. A slight modification to the original UNet, i.e., implementing dynamic weights on the vanilla skip connection, allows for a reasonable allocation of the fitting overhead of the loss function at different , and can improve DPM’s performance simply and effectively. As shown in Fig. 3, we implement dynamic skip connection by a simple linear dynamic weight . Thus, the novel dynamic skip connection can be denoted as
where , , and refer to the input, output, and intermediate layers, respectively. Furthermore, is an equation where is a pre-set hyperparameter that we choose as either 0.25 or 0.75. Notably, we replace all vanilla skip connections in UNet with dynamic skip connections.
Experiment
We conduct experiments to evaluate the effectiveness of CUD on CIFAR-10 with a resolution of 3232, excluding the comparison experiments. We use three datasets for the comparison experiments: CIFAR-10 with resolution 3232, MNIST with resolution 2828, and ImageNet-64 with resolution 6464. We assess the quality of the synthetic samples using Fréchet Inception Distance (FID) and Inception Score (IS) . To compute the FID, we compare 50,000 synthetic samples with all available real samples, and for computing the IS, we used 50,000 synthetic samples. We employed 4 different configurations for the velocity estimation model and hyperparameters on CIFAR-10, MNIST, and ImageNet-64. Specifically, we used configurations (a), (b), and (c) to train DPMs on CIFAR-10, MNIST, and ImageNet-64, respectively, and configuration (d) for our proposed final multi-step distillation (see Appendix G) on all datasets. The implementation details for these settings are provided in Appendix J. Unless otherwise specified, configuration (a) was used for ablation studies and analyses.
To verify that the basic CUD is capable of accelerated sampling, we conduct experiments and present the results in Table 1. In this table, the dynamic weight represents the supervision focus during training. At the beginning of the training, dynamic weight emphasizes using the ground truth label for supervision. In contrast, at the end of the training, dynamic weight emphasizes using the output of the teacher model for supervision. And is one of Learned Perceptual Image Patch Similarity (LPIPS) and MSE. LPIPS is more effective than MSE in Consistency Model . But according to Table 1, we can conclude that MSE can work, but LPIPS does not, and that vanilla weight can work but dynamic weight does not. Meanwhile, the effective choice of loss function , as well as the loss weights and in the basic CUD, are MSE, , and , respectively.
Investigate Design Space.
We present experimental results in Figs. 4 and 5 to assess the efficacy of our proposed suitable strategies. In these figures, Runge-Kutta 12/23 represents the form utilized in Runge-Kutta-based multi-step alignment distillation, while and denote the employment of the EMA model and the training model for catch-up sampling, respectively. Here, we do not consider the ground truth loss terms and . In Appendix I, we perform ablation experiments to show their importance. For all intermediate sampling points, we employ velocity estimation using a one-step method to synthetic samples. Specifically, when obtaining an intermediate sampling point , we calculate the target “clean” image using (Euler’s method). From Fig. 4, we deduce that employing the training model for catch-up sampling is more beneficial than using the EMA model, suggesting the application of instead of in Eq. 4. CUD is more effective in scenarios with fewer sampling steps but may be less effective than the baseline when more sampling steps are present. This is because CUD’s primary purpose is to accelerate sampling rather than improve the quality of synthetic images. Consequently, the velocity estimation model should prioritize generating “clean” images in scenarios with fewer sampling steps rather than enhancing image quality in scenarios with more sampling steps. Additionally, in Fig. 5, we observe that both uniform and rule surpass the baseline in terms of synthetic image quality in nearly all scenarios, with greater improvements seen with fewer sampling steps. In fewer sampling steps scenarios, rule outperforms uniform, while the opposite is true for scenarios with more sampling steps (details are provided in Appendix K). Notably, uniform achieves the best CUD performance, with an FID of 3.36. Lastly, the ablation experiments on dynamic skip connections are also presented in Figs. 5. Both 0.25 and 0.75 work very well because the fitting cost is minimized when is close to 0.5. In particular, 0.75 performs better than 0.25, achieving an FID 2.91 in 15 steps of sampling (see Appendix K). This suggests that to improve the quality of the synthetic images, the model should focus more on enhancing the representation at rather than . By combining the above three suitable strategies that require no additional overhead and are simple yet effective, the best performance of CUD on FID has decreased from 5.20 (4 steps to sampling with Heun’s method) to 2.91 (15 steps to sampling with Euler’s method).
Runge-Kutta-Based Multi-Step Alignment Distillation.
We perform ablation studies to assess the efficacy of Runge-Kutta-based multi-step alignment distillation across different orders, specifically utilizing the CIFAR-10 and MNIST datasets. The corresponding results are displayed in Tables 3 and 3. When examining the performance of Runge-Kutta orders 12, 23, and 34 on CIFAR-10, it becomes evident that images produced by Runge-Kutta 23 yield the most desirable performance. However, Runge-Kutta 34 attains the greatest accelerated sampling effect, followed by Runge-Kutta 23 and 12. In contrast, for the MNIST dataset, Runge-Kutta 34 provides the best performance and the most effective accelerated sampling. These empirical findings suggest that the accelerated sampling effect intensifies with the increase in the order of Runge-Kutta-based multi-step alignment distillation. Although the quality of images generated by Runge-Kutta 34 on CIFAR-10 is not as optimal as those produced by Runge-Kutta 12 and 23, it is justified due to the higher complexity of CIFAR-10 compared to MNIST, and the greater emphasis placed by Runge-Kutta 34 on improving accelerated sampling rather than enhancing the quality of the synthetic images.
2 Comparison Experiments
The comparative experiments are performed on CIFAR-10, MNIST, and ImageNet-64 to highlight the performance benefits of CUD. The results of PD , CD , and Curvature algorithms on MNIST are derived from Rectified flow, as the results of these methods are available for analysis. The relevant experimental results are presented in Tables 3 and 3, and the remaining supplementary results are given in Appendix I. Considered as a one-session training approach, we first compare CUD with other one-session training methods, such as the train-free accelerated sampling techniques, namely “DPM-solver” and DDIM, as well as a variety of DPMs, specifically NCSN++, DDPM, Rectified Flow, and Curvature, which serves as the baseline. For experiments on all datasets, CUD is very effective in accelerating sampling. For instance, CUD achieves the best performance in the few-step image generation scenario. For instance, on CIFAR-10, a FID 4.70 obtained by a CUD (Runge-Kutta 34) sample in just 7 steps better than a FID 5.28 obtained by a DPM-solver-2 sample in 12 steps.
Moreover, for the two-session training scenario, since CUD is trained on a continuous time step, it suffers from performance limitations. This is largely attributed to the excessive number of time points required for fitting, making it slightly inferior to both PD and CD algorithms. To address this, we introduced a novel multi-step distillation algorithm that extends the distillation algorithm presented in , as a way to achieve one-step sampling. Specifically, the algorithm is designed to incrementally fit “clean” images obtained at different time steps (e.g., when t progresses from 1/2 to 5/16 to 1/8). The effectiveness of this approach has been validated through quantitative ablation experiments discussed in Appendix G. Leveraging this, we conduct distillation on the pre-trained model derived from CUD (Runge-Kutta 12), ultimately achieving a state-of-the-art FID of 3.37 by sampling in one step on CIFAR-10. Furthermore, our approach demonstrates cost-effectiveness in training, requiring only 620k (500k+120k) iterations on CIFAR-10 with a batch size of 128 and a model parameter number of 55M, compared to CD’s 2100k (1300k+800k) iterations with a batch size of 256 and a model parameter number of 62M. Performance-wise, on both MNIST and ImageNet-64, CUD demonstrates comparable results to PD and CD. Lastly, synthetic images generated by our approach can be found in Appendix L.
Conclusion
This paper presents Catch-Up Distillation (CUD), a method designed to integrate effortlessly with the existing diffusion model training paradigm, enabling high-quality image synthesis in fewer sampling steps, and eliminating the need for pre-training weights. The efficacy of CUD is substantiated through the validation of several datasets, including CIFAR-10, MNIST, and ImageNet-64. In subsequent research, we sincerely hope that our CUD can be extended to discrete time steps where only a small number of time points need to be fitted, thus further enhancing the generality of CUD.
References
Appendix B Consistency Distillation and Consistency Training
Pseudo-code and implementation details of Consistency Distillation (CD) and Consistency Training (CT) can be found in paper https://arxiv.org/abs/2303.01469 and github link https://github.com/openai/consistency_models. Although Song et. al. state that CD is capable of implementing the new state-of-the-art FID 3.55 under the condition of a single NFE, there still are some constraints on this:
CD and CT may not be effective for certain popular DPMs such as Rectified flow , NCSN++, and DDPM . When CD or CT is applied to Rectified flow under the continuous time steps scenarios, it causes the training to collapse in our experiments (Table 1). This means that CD/CT requires Karra’s method under discrete time steps scenarios for it to work properly.
Constraint II.
CT can only be applied in empirical PF ODE because if the model design for estimating “clean image”, i.e., , in the DPM is ignored, then it is equivalent to its not having the ground truth label for effective supervision. For CT in empirical PF ODE, it relies on to achieve implicit supervision. To be specific, CT’s core loss function can be rewritten as . Setting , and , the loss function can continue to be rewritten as . Since empirical PF ODE guarantees that is an identity function when The original paper is that when . However, from the same study is defined as . Therefore, our derivations align theoretically in both instances., i.e., , is applied to fit the ground truth label . If , is required to fit . Then, if , is required to fit … Therefore, through a chain reaction, if , is required to fit .
Of course, this exquisite form of chain constraint accomplishes effective supervision, but it also hinders its possible application to other DPMs.
Constraint III.
Although the CT algorithm is able to accomplish distillation without relying on pre-trained weights for accelerating sampling, its application limitations and not very impressive performance prevent it from being a general-purpose algorithm. Our proposed CUD exists to address this problem, and it is theoretically applicable to any continuous SDE/ODE algorithm w.r.t., for SDE Runge-Kutta-based multi-step alignment distillation can be done with DDIM .
Constraint IV.
The models utilized in CD/CT, possessing 62 million parameters, are larger than our proposed CUD, which encompasses 55 million parameters in the CIFAR-10 dataset. CD requires 80w iterations for training, not including the additional 130w iterations necessary to train the original diffusion model. In contrast, CUD, despite incorporating multi-step distillation, only necessitates a total of 62w iterations, broken down into 50w and 12w iterations, respectively.
Before explaining why our proposed base loss can work, we first need to explain why Consistency Distillation (CD) of Consistency Model cannot work. As you know, the loss function of CD can be denoted as
Suppose that is a norm, i.e., . Then we can continue to derive Eq. 8 as
satisfy , we have
The right half of the above equation contains three terms: , , and . Only these terms are small enough to ensure the stability of the distillation. CD ignores the difference between the EMA model and the training model, which causes the training to collapse extremely easily under continuous time steps scenarios, especially if the EMA model is not updated accurately at one step.
If changing the pre-trained weight to , the left-hand side of Eq. 10 can be used directly to complete the distillation process. This form is different from the standard DPM training and is already widely applied in other applications to a certain extent. In image classification tasks, distillation algorithms typically involve two core losses: (1) Cross-Entropy loss between the student model output and the ground truth label, and (2) Kullback-Leibler Divergence between the student model output and the teacher model output. Collaborative supervision between the teacher model output and the ground truth label significantly enhances the generalization ability of the student model. This approach can also be applied to the DPM for distillation without teacher weights, thus achieving better performance of the student model.
Therefore, a distance function in the form of a norm as a loss function and the application of the ground truth label for supervision is necessary for the DPM. Only by ensuring these two points, Eq. 10’s bound does not collapse into training by being too lenient.
Appendix D Derivation of Runge-Kutta-Based Multi-Step Alignment Distillation
For the ODE , the precision of estimating is crucial for performing catch-up sampling from . Specifically, the smaller the truncation error of the ODE solver, the more precise the numerical integration will be in obtaining . Runge-Kutta methods, which include Euler’s method and Heun’s methods, enable a consistent view of modeling catch-up sampling. In particular, Euler’s method is a first-order ODE solver with truncation error , and Heun’s method is a second-order ODE solver with truncation error . Euler’s method is the same as Runge-Kutta 12 but Heun’s method is a subset of Runge-Kutta 23, due to the number of constraint equations of Runge-Kutta 23 being less than the number of solution factors. We only consider Runge-Kutta 12, Runge-Kutta 23, and Runge-Kutta 34 in this work because of the additional computational overhead required for higher orders Runge-Kutta methods. Let us define catch-up sampling as , where is the update function of a one-step ODE solver that takes , and as inputs to estimate the integration of velocities . Based on this, Runge-Kutta methods can be modeled out as
As shown in Fig. 1, our objective is to sample several points and utilize Runge-Kutta methods to compute the solutions of ODE at multiple instances starting from these points. Subsequently, we aim to estimate the velocities from the obtained solutions and align them with the outputs of the various heads . This means that we need to get the estimate of and expand them by Taylor’s Theorem in , then the formula can be written as
This rewrite is the result of transferring the effect of to . In Eq. 15, we have a necessary constraint that , the factors in must remain fixed, since we need to ensure that all sampling points remain fixed. Otherwise, the computational overhead would increase exponentially. Then, we can obtain the following novel constraints:
This means that we only need to satisfy the above constraint, and then multi-step alignment distillation based on Runge-Kutta 23 and Runge-Kutta 34 can be achieved. A natural form of sampling is equidistant sampling, i.e., . Based on this, all the factors can be determined:
We can derive the following Runge-Kutta-based multi-step alignment distillation algorithm:
Appendix E Higher-Order Runge-Kutta-Based Multi-step Alignment Distillation Algorithm
Due to space limitations in the main paper, we present here procedures of the higher order Runge-Kutta-based multi-step alignment distillation approaches, i.e., Runge-Kutta 23 (Algorithm C) and Runge-Kutta 34 (Algorithm C).
Appendix F Cost-of-Fit Analysis for Scenarios with Different t𝑡t
The optimization objective of DPMs is very different from the usual deep learning optimization objective, and one of the key differences is that DPMs has an additional input: the time point . Ideally, the velocity estimation models corresponding to the different time points can form a set of functions: . The number of neural network parameters is so large that most studies today apply weight sharing for velocity estimation models at different time steps, thus reducing storage costs. This approach also poses a problem at another level, in that at the end of the training phase, the expectation of training losses of the velocity estimation model at different time points varies considerably, as illustrated in Fig. 6. Specifically, the expectation of the training loss will be greater as or , and smaller as . This result is in line with our expectations, because as , the input is essentially free of Gaussian noise and can be approximated as , but the output has to be predicted as , when is completely unknowable and therefore extremely difficult. Similarly, as , the input is essentially Gaussian noise and can be approximated as , but the output has to be predicted as , when is completely unknowable and therefore also equally difficult.
Appendix G Final Multi-Step Distillation
Although CUD can accelerate sampling within a single training session, its reliance on continuous time steps necessitates fitting more time points than discrete time steps to complete the training, resulting in inferior performance compared to CD. In our experiments, merely replacing continuous time steps with discrete ones causes the model to collapse, as illustrated in Fig. 7. Thus, we adopt the distillation technique (a.k.a., one-step distillation) from to fit the velocity estimation model under a single time point, enabling the model to outperform CD in a one-step sampling scenario. Importantly, we incorporate ideas from the teacher-assistant concept to enhance the original distillation technique . Viewing the sample obtained in the last sampling step as a strong teacher output, the penultimate and penultimate third steps can be considered outputs from slightly weaker teachers. We can allow the model to gradually transition from an alignment time step from to , effectively preventing performance degradation due to the gap between the strong teacher and the weak student.
For instance, if we have sampled 16 steps using Euler’s method, we obtain a set of “clean” images for each sampling time point. We can first distill the model based on the “clean” images obtained in step 8, followed by those from steps 11 and 14. To ensure the model acquires a flat loss landscape after training, we align the training model’s output with the EMA model’s output, an approach interpretable as Sharpness-Aware Minimization (SAM) . In final multi-step distillation, aligning the outputs of the training model and the EMA model does not always yield successful results. Thereby, we regard this approach as an optional method.
In our experiments, we employ final multi-step distillation using “clean” images obtained at steps 8, 11, and 14 with a sampling of 16 steps through Euler’s method. To confirm the enhanced efficiency of our method compared to one-step distillation, we conduct comparative experiments and present the results in Table 4. Final multi-step distillation demonstrates a significant improvement in FID evaluation compared to one-step distillation, indicating the effectiveness and reasonableness of the multi-stage guided distillation. The positive impact of SAM is less evident and only boosts IS without affecting FID. Furthermore, the effectiveness of each step in final multi-step distillation is illustrated in Table 5. By guiding the training model through distillation using “clean” images from different steps, the synthetic images produced by the training model progressively improve in quality.
Appendix H Discussion
The CUD algorithm, which requires no pre-trained weights and a single training session, has considerably enhanced sampling acceleration compared to traditional DPM training paradigms. However, when assessed on the FID, a discernible performance gap with two-stage distillation-based accelerated sampling algorithms persists. This shortfall arises from two main factors: (1) Unlike those two-stage counterparts, which are trained on discrete time steps and necessitate a minimal number of time points to be fitted, CUD optimizes the loss function based on continuous time steps, demanding an infinite number of time points to be fitted. Yet, applying CUD directly to discrete time steps, as illustrated in Fig. 7, leads to training collapse due to instability. (2) CUD’s training duration is significantly shorter than that of two-stage distillation-based accelerated sampling algorithms.
Addressing this issue is a future work. This will involve revisiting hyperparameter selection and incorporating regularization terms to enhance training stability, thus enabling CUD to maintain stability over discrete time steps.
Broader Impact.
The idea of CUD is straightforward. It posits that the velocity estimation model output should align with the output from the same model at the previous moment, while simultaneously maintaining alignment with the ground truth label. This principle is anticipated to find broad application in future DPM training, potentially emerging as a more generalized paradigm. Consequently, we posit that CUD will exert a predominantly positive, rather than negative, impact on the community.
Appendix I Additional Experimental Results
In this section, we present a series of experimental results that can not fit in the main paper due to space constraints, including Tables 6, 7, 9, 9, 10, and Fig. 8. Tables 6 and 7 supplement the ablation experiments for Basic CUD in the main paper, while Tables 9, 9, and 10 supplement the comparison experimental results in the main paper. In particular, Tables 9, 9 are evaluated over samples sampled by the classical solver, including Euler and Heun, and Table 10 is evaluated over samples sampled by training-free accelerated sampling solvers, including DPM-solver-2 and Deis-2 .
We investigate the necessity of ground truth label supervision in this paragraph, showcasing both the Runge-Kutta-based multi-step alignment distillation that incorporates ground truth label supervision and its counterpart without such supervision in Fig. 8. The label “Runge-Kutta 23 (w/o GT)” signifies CUD lacking the loss term , while “Runge-Kutta 23 (with GT)” denotes the same algorithm but includes this loss term. Similarly, “Runge-Kutta 34 (w/o GT)” indicates CUD without the loss terms and , and “Runge-Kutta 23 (with GT)” implies the inclusion of these loss terms. The ground truth supervision is essential for CUD, as it helps circumvent improper gradient updates that could arise from unsupervised distillation and thus adversely affect performance. As a result, for in the main paper, we incorporate ground truth supervision into each head of the velocity estimation model .
Appendix J Implementation Details
Table. 11 shows the training and architecture configuration we use in our experiments. For CIFAR-10, MNIST, and ImageNet-64 datasets, we carry out three different configurations (i.e., (a), (b), and (c)) for experiments, where configurations (a) and (b) are following and the model architecture in configuration (c) is following . We ran the code on NVIDIA Tesla A100 GPUs, where configurations (a) used 4 GPUs with a batch size of 32 on each GPU. Configuration (b) applied 4 GPUs with a batch size of 64 on each GPU. Configuration (c) applied 8 GPUs with a batch size of 128 on each GPU. Configuration (d) is used for the final multi-step distillation, and its batch size and the learning rate are strongly correlated with the dataset. For example, if the final discrete distillation is performed on MNIST, then its batch size and learning rate are the same as configuration (b).
Appendix K Additional Content Presentation
When , the comparison between the curves in Figs. 4 and 5 in the main paper are not fuzzy. Therefore, we present the specific data in Table 12 and 13. These tables are labeled in exactly the same form and with the same data content as Figs. 4 and 5.
Appendix L Additional Samples from Catch-Up Distillation and Final Multi-Step Distillation
We provide additional samples from Catch-Up Distillation (CUD) and Final Multi-Step Distillation (FMSD) on MNIST (Figs. 9), CIFAR-10 (Figs. 10), and ImageNet-64 (Figs. 11).