MCPcopy Create free account
hub / github.com/Zhiyuan-R/Tiger-Diffusion / p_sample

Method p_sample

test_generation.py:187–202  ·  view source on GitHub ↗

Sample from the model

(self, denoise_fn, data, t, noise_fn, clip_denoised=False, return_pred_xstart=False, use_var=True)

Source from the content-addressed store, hash-verified

185 ''' samples '''
186
187 def p_sample(self, denoise_fn, data, t, noise_fn, clip_denoised=False, return_pred_xstart=False, use_var=True):
188 """
189 Sample from the model
190 """
191 model_mean, _, model_log_variance, pred_xstart = self.p_mean_variance(denoise_fn, data=data, t=t, clip_denoised=clip_denoised,
192 return_pred_xstart=True)
193 noise = noise_fn(size=data.shape, dtype=data.dtype, device=data.device)
194 assert noise.shape == data.shape
195 # no noise when t == 0
196 nonzero_mask = torch.reshape(1 - (t == 0).float(), [data.shape[0]] + [1] * (len(data.shape) - 1))
197
198 sample = model_mean
199 if use_var:
200 sample = sample + nonzero_mask * torch.exp(0.5 * model_log_variance) * noise
201 assert sample.shape == pred_xstart.shape
202 return (sample, pred_xstart) if return_pred_xstart else sample
203
204
205 def p_sample_loop(self, denoise_fn, shape, device,

Callers 2

p_sample_loopMethod · 0.95
reconstructMethod · 0.95

Calls 1

p_mean_varianceMethod · 0.95

Tested by

no test coverage detected