| 375 | |
| 376 | @register_sampler(name='ddim') |
| 377 | class 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 | # ================= |
nothing calls this directly
no outgoing calls
no test coverage detected