Create a wrapper function for the noise prediction model. DPM-Solver needs to solve the continuous-time diffusion ODEs. For DPMs trained on discrete-time labels, we need to firstly wrap the model function to a noise prediction model that accepts the continuous time as the input. We suppo
(
model,
noise_schedule,
model_type="noise",
model_kwargs={},
guidance_type="uncond",
condition=None,
unconditional_condition=None,
guidance_scale=1.,
classifier_fn=None,
classifier_kwargs={},
)
| 159 | |
| 160 | |
| 161 | def model_wrapper( |
| 162 | model, |
| 163 | noise_schedule, |
| 164 | model_type="noise", |
| 165 | model_kwargs={}, |
| 166 | guidance_type="uncond", |
| 167 | condition=None, |
| 168 | unconditional_condition=None, |
| 169 | guidance_scale=1., |
| 170 | classifier_fn=None, |
| 171 | classifier_kwargs={}, |
| 172 | ): |
| 173 | """Create a wrapper function for the noise prediction model. |
| 174 | DPM-Solver needs to solve the continuous-time diffusion ODEs. For DPMs trained on discrete-time labels, we need to |
| 175 | firstly wrap the model function to a noise prediction model that accepts the continuous time as the input. |
| 176 | We support four types of the diffusion model by setting `model_type`: |
| 177 | 1. "noise": noise prediction model. (Trained by predicting noise). |
| 178 | 2. "x_start": data prediction model. (Trained by predicting the data x_0 at time 0). |
| 179 | 3. "v": velocity prediction model. (Trained by predicting the velocity). |
| 180 | The "v" prediction is derivation detailed in Appendix D of [1], and is used in Imagen-Video [2]. |
| 181 | [1] Salimans, Tim, and Jonathan Ho. "Progressive distillation for fast sampling of diffusion models." |
| 182 | arXiv preprint arXiv:2202.00512 (2022). |
| 183 | [2] Ho, Jonathan, et al. "Imagen Video: High Definition Video Generation with Diffusion Models." |
| 184 | arXiv preprint arXiv:2210.02303 (2022). |
| 185 | |
| 186 | 4. "score": marginal score function. (Trained by denoising score matching). |
| 187 | Note that the score function and the noise prediction model follows a simple relationship: |
| 188 | ``` |
| 189 | noise(x_t, t) = -sigma_t * score(x_t, t) |
| 190 | ``` |
| 191 | We support three types of guided sampling by DPMs by setting `guidance_type`: |
| 192 | 1. "uncond": unconditional sampling by DPMs. |
| 193 | The input `model` has the following format: |
| 194 | `` |
| 195 | model(x, t_input, **model_kwargs) -> noise | x_start | v | score |
| 196 | `` |
| 197 | 2. "classifier": classifier guidance sampling [3] by DPMs and another classifier. |
| 198 | The input `model` has the following format: |
| 199 | `` |
| 200 | model(x, t_input, **model_kwargs) -> noise | x_start | v | score |
| 201 | `` |
| 202 | The input `classifier_fn` has the following format: |
| 203 | `` |
| 204 | classifier_fn(x, t_input, cond, **classifier_kwargs) -> logits(x, t_input, cond) |
| 205 | `` |
| 206 | [3] P. Dhariwal and A. Q. Nichol, "Diffusion models beat GANs on image synthesis," |
| 207 | in Advances in Neural Information Processing Systems, vol. 34, 2021, pp. 8780-8794. |
| 208 | 3. "classifier-free": classifier-free guidance sampling by conditional DPMs. |
| 209 | The input `model` has the following format: |
| 210 | `` |
| 211 | model(x, t_input, cond, **model_kwargs) -> noise | x_start | v | score |
| 212 | `` |
| 213 | And if cond == `unconditional_condition`, the model output is the unconditional DPM output. |
| 214 | [4] Ho, Jonathan, and Tim Salimans. "Classifier-free diffusion guidance." |
| 215 | arXiv preprint arXiv:2207.12598 (2022). |
| 216 | |
| 217 | The `t_input` is the time label of the model, which may be discrete-time labels (i.e. 0 to 999) |
| 218 | or continuous-time labels (i.e. epsilon to T). |