MCPcopy Create free account
hub / github.com/TencentARC/MotionCtrl / p_sample

Method p_sample

lvdm/models/ddpm3d.py:862–881  ·  view source on GitHub ↗
(self, x, c, t, clip_denoised=False, repeat_noise=False, return_x0=False, \
                 temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, **kwargs)

Source from the content-addressed store, hash-verified

860
861 @torch.no_grad()
862 def p_sample(self, x, c, t, clip_denoised=False, repeat_noise=False, return_x0=False, \
863 temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, **kwargs):
864 b, *_, device = *x.shape, x.device
865 outputs = self.p_mean_variance(x=x, c=c, t=t, clip_denoised=clip_denoised, return_x0=return_x0, \
866 score_corrector=score_corrector, corrector_kwargs=corrector_kwargs, **kwargs)
867 if return_x0:
868 model_mean, _, model_log_variance, x0 = outputs
869 else:
870 model_mean, _, model_log_variance = outputs
871
872 noise = noise_like(x.shape, device, repeat_noise) * temperature
873 if noise_dropout > 0.:
874 noise = torch.nn.functional.dropout(noise, p=noise_dropout)
875 # no noise when t == 0
876 nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1)))
877
878 if return_x0:
879 return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise, x0
880 else:
881 return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise
882
883 @torch.no_grad()
884 def p_sample_loop(self, cond, shape, return_intermediates=False, x_T=None, verbose=True, callback=None, \

Callers 1

p_sample_loopMethod · 0.95

Calls 2

p_mean_varianceMethod · 0.95
noise_likeFunction · 0.90

Tested by

no test coverage detected