MCPcopy Create free account
hub / github.com/MegaScenes/nvs / p_sample_loop

Method p_sample_loop

ldm/models/diffusion/all_functions_ddpm.py:1247–1295  ·  view source on GitHub ↗
(self, cond, shape, return_intermediates=False,
                      x_T=None, verbose=True, callback=None, timesteps=None, quantize_denoised=False,
                      mask=None, x0=None, img_callback=None, start_T=None,
                      log_every_t=None)

Source from the content-addressed store, hash-verified

1245
1246 @torch.no_grad()
1247 def p_sample_loop(self, cond, shape, return_intermediates=False,
1248 x_T=None, verbose=True, callback=None, timesteps=None, quantize_denoised=False,
1249 mask=None, x0=None, img_callback=None, start_T=None,
1250 log_every_t=None):
1251
1252 if not log_every_t:
1253 log_every_t = self.log_every_t
1254 device = self.betas.device
1255 b = shape[0]
1256 if x_T is None:
1257 img = torch.randn(shape, device=device)
1258 else:
1259 img = x_T
1260
1261 intermediates = [img]
1262 if timesteps is None:
1263 timesteps = self.num_timesteps
1264
1265 if start_T is not None:
1266 timesteps = min(timesteps, start_T)
1267 iterator = tqdm(reversed(range(0, timesteps)), desc='Sampling t', total=timesteps) if verbose else reversed(
1268 range(0, timesteps))
1269
1270 if mask is not None:
1271 assert x0 is not None
1272 assert x0.shape[2:3] == mask.shape[2:3] # spatial size has to match
1273
1274 for i in iterator:
1275 ts = torch.full((b,), i, device=device, dtype=torch.long)
1276 if self.shorten_cond_schedule:
1277 assert self.model.conditioning_key != 'hybrid'
1278 tc = self.cond_ids[ts].to(cond.device)
1279 cond = self.q_sample(x_start=cond, t=tc, noise=torch.randn_like(cond))
1280
1281 img = self.p_sample(img, cond, ts,
1282 clip_denoised=self.clip_denoised,
1283 quantize_denoised=quantize_denoised)
1284 if mask is not None:
1285 img_orig = self.q_sample(x0, ts)
1286 img = img_orig * mask + (1. - mask) * img
1287
1288 if i % log_every_t == 0 or i == timesteps - 1:
1289 intermediates.append(img)
1290 if callback: callback(i)
1291 if img_callback: img_callback(img, i)
1292
1293 if return_intermediates:
1294 return img, intermediates
1295 return img
1296
1297 @torch.no_grad()
1298 def sample(self, cond, batch_size=16, return_intermediates=False, x_T=None,

Callers 1

sampleMethod · 0.95

Calls 3

p_sampleMethod · 0.95
toMethod · 0.80
q_sampleMethod · 0.45

Tested by

no test coverage detected