(model, x, sigma_min, sigma_max, extra_args=None, atol=1e-4, rtol=1e-4)
| 279 | |
| 280 | @torch.no_grad() |
| 281 | def log_likelihood(model, x, sigma_min, sigma_max, extra_args=None, atol=1e-4, rtol=1e-4): |
| 282 | extra_args = {} if extra_args is None else extra_args |
| 283 | s_in = x.new_ones([x.shape[0]]) |
| 284 | v = torch.randint_like(x, 2) * 2 - 1 |
| 285 | fevals = 0 |
| 286 | def ode_fn(sigma, x): |
| 287 | nonlocal fevals |
| 288 | with torch.enable_grad(): |
| 289 | x = x[0].detach().requires_grad_() |
| 290 | denoised = model(x, sigma * s_in, **extra_args) |
| 291 | d = to_d(x, sigma, denoised) |
| 292 | fevals += 1 |
| 293 | grad = torch.autograd.grad((d * v).sum(), x)[0] |
| 294 | d_ll = (v * grad).flatten(1).sum(1) |
| 295 | return d.detach(), d_ll |
| 296 | x_min = x, x.new_zeros([x.shape[0]]) |
| 297 | t = x.new_tensor([sigma_min, sigma_max]) |
| 298 | sol = odeint(ode_fn, x_min, t, atol=atol, rtol=rtol, method='dopri5') |
| 299 | latent, delta_ll = sol[0][-1], sol[1][-1] |
| 300 | ll_prior = torch.distributions.Normal(0, sigma_max).log_prob(latent).flatten(1).sum(1) |
| 301 | return ll_prior + delta_ll, {'fevals': fevals} |
| 302 | |
| 303 | |
| 304 | class PIDStepSizeController: |
nothing calls this directly
no outgoing calls
no test coverage detected