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,
pre_seq=None,
)
| 771 | img = out["sample"] |
| 772 | |
| 773 | def ddim_sample( |
| 774 | self, |
| 775 | model, |
| 776 | x, |
| 777 | t, |
| 778 | clip_denoised=True, |
| 779 | denoised_fn=None, |
| 780 | cond_fn=None, |
| 781 | model_kwargs=None, |
| 782 | eta=0.0, |
| 783 | pre_seq=None, |
| 784 | ): |
| 785 | """ |
| 786 | Sample x_{t-1} from the model using DDIM. |
| 787 | |
| 788 | Same usage as p_sample(). |
| 789 | """ |
| 790 | if pre_seq is not None: |
| 791 | T = pre_seq.shape[1] |
| 792 | noise = th.randn_like(pre_seq) |
| 793 | x_t = self.q_sample(pre_seq, t, noise=noise) |
| 794 | x[:, :T, :] = x_t |
| 795 | |
| 796 | out = self.p_mean_variance( |
| 797 | model, |
| 798 | x, |
| 799 | t, |
| 800 | clip_denoised=clip_denoised, |
| 801 | denoised_fn=denoised_fn, |
| 802 | model_kwargs=model_kwargs, |
| 803 | ) |
| 804 | if cond_fn is not None: |
| 805 | out = self.condition_score(cond_fn, |
| 806 | out, |
| 807 | x, |
| 808 | t, |
| 809 | model_kwargs=model_kwargs) |
| 810 | |
| 811 | # Usually our model outputs epsilon, but we re-derive it |
| 812 | # in case we used x_start or x_prev prediction. |
| 813 | eps = self._predict_eps_from_xstart(x, t, out["pred_xstart"]) |
| 814 | |
| 815 | alpha_bar = _extract_into_tensor(self.alphas_cumprod, t, x.shape) |
| 816 | alpha_bar_prev = _extract_into_tensor(self.alphas_cumprod_prev, t, |
| 817 | x.shape) |
| 818 | sigma = (eta * th.sqrt((1 - alpha_bar_prev) / (1 - alpha_bar)) * |
| 819 | th.sqrt(1 - alpha_bar / alpha_bar_prev)) |
| 820 | # Equation 12. |
| 821 | noise = th.randn_like(x) |
| 822 | mean_pred = (out["pred_xstart"] * th.sqrt(alpha_bar_prev) + |
| 823 | th.sqrt(1 - alpha_bar_prev - sigma**2) * eps) |
| 824 | nonzero_mask = ((t != 0).float().view(-1, *([1] * (len(x.shape) - 1))) |
| 825 | ) # no noise when t == 0 |
| 826 | sample = mean_pred + nonzero_mask * sigma * noise |
| 827 | return {"sample": sample, "pred_xstart": out["pred_xstart"]} |
| 828 | |
| 829 | def ddim_reverse_sample( |
| 830 | self, |
no test coverage detected