MCPcopy Create free account
hub / github.com/Sin3DM/Sin3DM / ddim_sample

Method ddim_sample

src/diffusion/gaussian_diffusion.py:538–600  ·  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,
        y0=None,
        mask=None,
        is_mask_t0=False,
    )

Source from the content-addressed store, hash-verified

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 )

Callers 1

Calls 4

p_mean_varianceMethod · 0.95
condition_scoreMethod · 0.95
_extract_into_tensorFunction · 0.85

Tested by

no test coverage detected