MCPcopy Create free account
hub / github.com/VisionXLab/OF-Diff / model_wrapper

Function model_wrapper

ldm/models/diffusion/dpm_solver/dpm_solver.py:161–316  ·  view source on GitHub ↗

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={},
)

Source from the content-addressed store, hash-verified

159
160
161def 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).

Callers 1

sampleMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected