The noise predicition model function that is used for DPM-Solver.
(x, t_continuous)
| 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"] |
nothing calls this directly
no test coverage detected