MCPcopy Create free account
hub / github.com/MotrixLab/FineMoGen / ddim_sample

Method ddim_sample

mogen/models/utils/gaussian_diffusion.py:773–827  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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,

Callers 1

Calls 5

q_sampleMethod · 0.95
p_mean_varianceMethod · 0.95
condition_scoreMethod · 0.95
_extract_into_tensorFunction · 0.85

Tested by

no test coverage detected