(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)
| 1264 | |
| 1265 | @torch.no_grad() |
| 1266 | def sample(self, cond, batch_size=16, return_intermediates=False, x_T=None, |
| 1267 | verbose=True, timesteps=None, quantize_denoised=False, |
| 1268 | mask=None, x0=None, shape=None,**kwargs): |
| 1269 | if shape is None: |
| 1270 | shape = (batch_size, self.channels, self.image_size, self.image_size) |
| 1271 | if cond is not None: |
| 1272 | if isinstance(cond, dict): |
| 1273 | cond = {key: cond[key][:batch_size] if not isinstance(cond[key], list) else |
| 1274 | list(map(lambda x: x[:batch_size], cond[key])) for key in cond} |
| 1275 | else: |
| 1276 | cond = [c[:batch_size] for c in cond] if isinstance(cond, list) else cond[:batch_size] |
| 1277 | return self.p_sample_loop(cond, |
| 1278 | shape, |
| 1279 | return_intermediates=return_intermediates, x_T=x_T, |
| 1280 | verbose=verbose, timesteps=timesteps, quantize_denoised=quantize_denoised, |
| 1281 | mask=mask, x0=x0) |
| 1282 | |
| 1283 | @torch.no_grad() |
| 1284 | def sample_log(self,cond,batch_size,ddim, ddim_steps,**kwargs): |
no test coverage detected