(self, cond, batch_size=16, return_intermediates=False, x_T=None,
verbose=True, timesteps=None, quantize_denoised=False,
mask=None, x0=None, shape=None,**kwargs)
| 1296 | |
| 1297 | @torch.no_grad() |
| 1298 | def sample(self, cond, batch_size=16, return_intermediates=False, x_T=None, |
| 1299 | verbose=True, timesteps=None, quantize_denoised=False, |
| 1300 | mask=None, x0=None, shape=None,**kwargs): |
| 1301 | if shape is None: |
| 1302 | shape = (batch_size, self.channels, self.image_size, self.image_size) |
| 1303 | if cond is not None: |
| 1304 | if isinstance(cond, dict): |
| 1305 | cond = {key: cond[key][:batch_size] if not isinstance(cond[key], list) else |
| 1306 | list(map(lambda x: x[:batch_size], cond[key])) for key in cond} |
| 1307 | else: |
| 1308 | cond = [c[:batch_size] for c in cond] if isinstance(cond, list) else cond[:batch_size] |
| 1309 | return self.p_sample_loop(cond, |
| 1310 | shape, |
| 1311 | return_intermediates=return_intermediates, x_T=x_T, |
| 1312 | verbose=verbose, timesteps=timesteps, quantize_denoised=quantize_denoised, |
| 1313 | mask=mask, x0=x0) |
| 1314 | |
| 1315 | @torch.no_grad() |
| 1316 | def sample_log(self, cond, batch_size, ddim, ddim_steps, **kwargs): |
no test coverage detected