MCPcopy Create free account
hub / github.com/openai/point-e / ddim_sample

Method ddim_sample

point_e/diffusion/gaussian_diffusion.py:550–598  ·  view source on GitHub ↗

Sample x_{t-1} from the model using DDIM. Same usage as p_sample().

(
        self,
        model,
        x,
        t,
        clip_denoised=False,
        denoised_fn=None,
        cond_fn=None,
        model_kwargs=None,
        eta=0.0,
    )

Source from the content-addressed store, hash-verified

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

Callers 1

Calls 4

p_mean_varianceMethod · 0.95
condition_scoreMethod · 0.95
_extract_into_tensorFunction · 0.85

Tested by

no test coverage detected