Sample x_{t-1} from the model using DDIM. Same usage as p_sample().
(
self,
model,
x,
t,
clip_denoised=True,
denoised_fn=None,
cond_fn=None,
model_kwargs=None,
eta=0.0,
y0=None,
mask=None,
is_mask_t0=False,
)
| 536 | img = out["sample"] |
| 537 | |
| 538 | def ddim_sample( |
| 539 | self, |
| 540 | model, |
| 541 | x, |
| 542 | t, |
| 543 | clip_denoised=True, |
| 544 | denoised_fn=None, |
| 545 | cond_fn=None, |
| 546 | model_kwargs=None, |
| 547 | eta=0.0, |
| 548 | y0=None, |
| 549 | mask=None, |
| 550 | is_mask_t0=False, |
| 551 | ): |
| 552 | """ |
| 553 | Sample x_{t-1} from the model using DDIM. |
| 554 | |
| 555 | Same usage as p_sample(). |
| 556 | """ |
| 557 | out = self.p_mean_variance( |
| 558 | model, |
| 559 | x, |
| 560 | t, |
| 561 | clip_denoised=clip_denoised, |
| 562 | denoised_fn=denoised_fn, |
| 563 | model_kwargs=model_kwargs, |
| 564 | ) |
| 565 | if cond_fn is not None: |
| 566 | out = self.condition_score(cond_fn, out, x, t, model_kwargs=model_kwargs) |
| 567 | # masked generation |
| 568 | if y0 is not None and mask is not None: |
| 569 | assert y0.shape == x.shape |
| 570 | assert mask.shape == x.shape |
| 571 | if is_mask_t0: |
| 572 | out["pred_xstart"] = mask * y0 + (1 - mask) * out["pred_xstart"] |
| 573 | else: |
| 574 | nonzero_mask = ( |
| 575 | (t != 0).float().view(-1, *([1] * (len(x.shape) - 1))) |
| 576 | ) # no noise when t == 0 |
| 577 | out["pred_xstart"] = (mask * y0 + (1 - mask) * out["pred_xstart"]) * nonzero_mask + out["pred_xstart"] * (1 - nonzero_mask) |
| 578 | |
| 579 | # Usually our model outputs epsilon, but we re-derive it |
| 580 | # in case we used x_start or x_prev prediction. |
| 581 | eps = self._predict_eps_from_xstart(x, t, out["pred_xstart"]) |
| 582 | |
| 583 | alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape) |
| 584 | alpha_bar_prev = _extract_into_tensor(self.alphas_cumprod_prev, t, x.shape) |
| 585 | sigma = ( |
| 586 | eta |
| 587 | * th.sqrt((1 - alpha_bar_prev) / (1 - alpha_bar)) |
| 588 | * th.sqrt(1 - alpha_bar / alpha_bar_prev) |
| 589 | ) |
| 590 | # Equation 12. |
| 591 | noise = th.randn_like(x) |
| 592 | mean_pred = ( |
| 593 | out["pred_xstart"] * th.sqrt(alpha_bar_prev) |
| 594 | + th.sqrt(1 - alpha_bar_prev - sigma ** 2) * eps |
| 595 | ) |
no test coverage detected