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

Function model_fn

ldm/models/diffusion/dpm_solver/dpm_solver.py:289–312  ·  view source on GitHub ↗

The noise predicition model function that is used for DPM-Solver.

(x, t_continuous)

Source from the content-addressed store, hash-verified

287 return torch.autograd.grad(log_prob.sum(), x_in)[0]
288
289 def model_fn(x, t_continuous):
290 """
291 The noise predicition model function that is used for DPM-Solver.
292 """
293 if t_continuous.reshape((-1,)).shape[0] == 1:
294 t_continuous = t_continuous.expand((x.shape[0]))
295 if guidance_type == "uncond":
296 return noise_pred_fn(x, t_continuous)
297 elif guidance_type == "classifier":
298 assert classifier_fn is not None
299 t_input = get_model_input_time(t_continuous)
300 cond_grad = cond_grad_fn(x, t_input)
301 sigma_t = noise_schedule.marginal_std(t_continuous)
302 noise = noise_pred_fn(x, t_continuous)
303 return noise - guidance_scale * expand_dims(sigma_t, dims=cond_grad.dim()) * cond_grad
304 elif guidance_type == "classifier-free":
305 if guidance_scale == 1. or unconditional_condition is None:
306 return noise_pred_fn(x, t_continuous, cond=condition)
307 else:
308 x_in = torch.cat([x] * 2)
309 t_in = torch.cat([t_continuous] * 2)
310 c_in = torch.cat([unconditional_condition, condition])
311 noise_uncond, noise = noise_pred_fn(x_in, t_in, cond=c_in).chunk(2)
312 return noise_uncond + guidance_scale * (noise - noise_uncond)
313
314 assert model_type in ["noise", "x_start", "v"]
315 assert guidance_type in ["uncond", "classifier", "classifier-free"]

Callers

nothing calls this directly

Calls 5

noise_pred_fnFunction · 0.85
get_model_input_timeFunction · 0.85
cond_grad_fnFunction · 0.85
expand_dimsFunction · 0.85
marginal_stdMethod · 0.80

Tested by

no test coverage detected