MCPcopy Create free account
hub / github.com/DPS2022/diffusion-posterior-sampling / DDIM

Class DDIM

guided_diffusion/gaussian_diffusion.py:377–406  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

375
376@register_sampler(name='ddim')
377class DDIM(SpacedDiffusion):
378 def p_sample(self, model, x, t, eta=0.0):
379 out = self.p_mean_variance(model, x, t)
380
381 eps = self.predict_eps_from_x_start(x, t, out['pred_xstart'])
382
383 alpha_bar = extract_and_expand(self.alphas_cumprod, t, x)
384 alpha_bar_prev = extract_and_expand(self.alphas_cumprod_prev, t, x)
385 sigma = (
386 eta
387 * torch.sqrt((1 - alpha_bar_prev) / (1 - alpha_bar))
388 * torch.sqrt(1 - alpha_bar / alpha_bar_prev)
389 )
390 # Equation 12.
391 noise = torch.randn_like(x)
392 mean_pred = (
393 out["pred_xstart"] * torch.sqrt(alpha_bar_prev)
394 + torch.sqrt(1 - alpha_bar_prev - sigma ** 2) * eps
395 )
396
397 sample = mean_pred
398 if t != 0:
399 sample += sigma * noise
400
401 return {"sample": sample, "pred_xstart": out["pred_xstart"]}
402
403 def predict_eps_from_x_start(self, x_t, t, pred_xstart):
404 coef1 = extract_and_expand(self.sqrt_recip_alphas_cumprod, t, x_t)
405 coef2 = extract_and_expand(self.sqrt_recipm1_alphas_cumprod, t, x_t)
406 return (coef1 * x_t - pred_xstart) / coef2
407
408
409# =================

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected