MCPcopy Create free account
hub / github.com/Meshcapade/difflocks / log_likelihood

Function log_likelihood

k_diffusion/sampling.py:281–301  ·  view source on GitHub ↗
(model, x, sigma_min, sigma_max, extra_args=None, atol=1e-4, rtol=1e-4)

Source from the content-addressed store, hash-verified

279
280@torch.no_grad()
281def 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
304class PIDStepSizeController:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected