MCPcopy Create free account
hub / github.com/IceClear/StableSR / p_sample_loop

Method p_sample_loop

ldm/models/diffusion/ddpm.py:1331–1379  ·  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

1329
1330 @torch.no_grad()
1331 def p_sample_loop(self, cond, shape, return_intermediates=False,
1332 x_T=None, verbose=True, callback=None, timesteps=None, quantize_denoised=False,
1333 mask=None, x0=None, img_callback=None, start_T=None,
1334 log_every_t=None):
1335
1336 if not log_every_t:
1337 log_every_t = self.log_every_t
1338 device = self.betas.device
1339 b = shape[0]
1340 if x_T is None:
1341 img = torch.randn(shape, device=device)
1342 else:
1343 img = x_T
1344
1345 intermediates = [img]
1346 if timesteps is None:
1347 timesteps = self.num_timesteps
1348
1349 if start_T is not None:
1350 timesteps = min(timesteps, start_T)
1351 iterator = tqdm(reversed(range(0, timesteps)), desc='Sampling t', total=timesteps) if verbose else reversed(
1352 range(0, timesteps))
1353
1354 if mask is not None:
1355 assert x0 is not None
1356 assert x0.shape[2:3] == mask.shape[2:3] # spatial size has to match
1357
1358 for i in iterator:
1359 ts = torch.full((b,), i, device=device, dtype=torch.long)
1360 if self.shorten_cond_schedule:
1361 assert self.model.conditioning_key != 'hybrid'
1362 tc = self.cond_ids[ts].to(cond.device)
1363 cond = self.q_sample(x_start=cond, t=tc, noise=torch.randn_like(cond))
1364
1365 img = self.p_sample(img, cond, ts,
1366 clip_denoised=self.clip_denoised,
1367 quantize_denoised=quantize_denoised)
1368 if mask is not None:
1369 img_orig = self.q_sample(x0, ts)
1370 img = img_orig * mask + (1. - mask) * img
1371
1372 if i % log_every_t == 0 or i == timesteps - 1:
1373 intermediates.append(img)
1374 if callback: callback(i)
1375 if img_callback: img_callback(img, i)
1376
1377 if return_intermediates:
1378 return img, intermediates
1379 return img
1380
1381 @torch.no_grad()
1382 def sample(self, cond, batch_size=16, return_intermediates=False, x_T=None,

Callers 1

sampleMethod · 0.95

Calls 3

p_sampleMethod · 0.95
callbackFunction · 0.85
q_sampleMethod · 0.45

Tested by

no test coverage detected