dagloop5 commited on
Commit
e2a5a9a
·
verified ·
1 Parent(s): 84c1ca9

Update pk_workflow.py

Browse files
Files changed (1) hide show
  1. pk_workflow.py +18 -0
pk_workflow.py CHANGED
@@ -284,6 +284,12 @@ def _er_sde_step(scheduler, generator, model_output, timestep, sample, s_noise:
284
 
285
  if s_noise > 0:
286
  noise = torch.randn(x.shape, dtype=x.dtype, device="cpu", generator=generator).to(x.device)
 
 
 
 
 
 
287
  spread = (er_lambda_t**2 - er_lambda_s**2 * r**2).clamp_min(0).sqrt()
288
  prev_sample = prev_sample + alpha_t * noise * s_noise * spread
289
 
@@ -401,6 +407,12 @@ def _dpmpp_2m_sde_step(scheduler, model_output, timestep, sample, eta: float = 1
401
 
402
  if eta > 0 and s_noise > 0:
403
  noise = scheduler._dpmpp_sde_noise_sampler(sigma, sigma_next).to(device=x.device, dtype=compute_dtype)
 
 
 
 
 
 
404
  prev_sample = prev_sample + noise * sigma_next * (-2 * h * eta).expm1().neg().sqrt() * s_noise
405
 
406
  scheduler._dpmpp_sde_h_last = h
@@ -492,6 +504,12 @@ def _dpmpp_3m_sde_step(scheduler, model_output, timestep, sample, eta: float = 1
492
 
493
  if eta > 0 and s_noise > 0:
494
  noise = scheduler._dpmpp_sde_noise_sampler(sigma, sigma_next).to(device=x.device, dtype=compute_dtype)
 
 
 
 
 
 
495
  prev_sample = prev_sample + noise * sigma_next * (-2 * h * eta).expm1().neg().sqrt() * s_noise
496
 
497
  scheduler._dpmpp_sde_h_last_2 = h_1
 
284
 
285
  if s_noise > 0:
286
  noise = torch.randn(x.shape, dtype=x.dtype, device="cpu", generator=generator).to(x.device)
287
+ print(
288
+ f"[noise-debug] er_sde i={scheduler._step_index} x.shape={tuple(x.shape)} "
289
+ f"noise.shape={tuple(noise.shape)} noise.std={noise.std().item():.4f} "
290
+ f"noise.mean={noise.mean().item():.4f}",
291
+ flush=True,
292
+ )
293
  spread = (er_lambda_t**2 - er_lambda_s**2 * r**2).clamp_min(0).sqrt()
294
  prev_sample = prev_sample + alpha_t * noise * s_noise * spread
295
 
 
407
 
408
  if eta > 0 and s_noise > 0:
409
  noise = scheduler._dpmpp_sde_noise_sampler(sigma, sigma_next).to(device=x.device, dtype=compute_dtype)
410
+ print(
411
+ f"[noise-debug] dpmpp_2m_sde i={scheduler._step_index} x.shape={tuple(x.shape)} "
412
+ f"noise.shape={tuple(noise.shape)} noise.std={noise.std().item():.4f} "
413
+ f"noise.mean={noise.mean().item():.4f}",
414
+ flush=True,
415
+ )
416
  prev_sample = prev_sample + noise * sigma_next * (-2 * h * eta).expm1().neg().sqrt() * s_noise
417
 
418
  scheduler._dpmpp_sde_h_last = h
 
504
 
505
  if eta > 0 and s_noise > 0:
506
  noise = scheduler._dpmpp_sde_noise_sampler(sigma, sigma_next).to(device=x.device, dtype=compute_dtype)
507
+ print(
508
+ f"[noise-debug] dpmpp_3m_sde i={scheduler._step_index} x.shape={tuple(x.shape)} "
509
+ f"noise.shape={tuple(noise.shape)} noise.std={noise.std().item():.4f} "
510
+ f"noise.mean={noise.mean().item():.4f}",
511
+ flush=True,
512
+ )
513
  prev_sample = prev_sample + noise * sigma_next * (-2 * h * eta).expm1().neg().sqrt() * s_noise
514
 
515
  scheduler._dpmpp_sde_h_last_2 = h_1