dagloop5 commited on
Commit
8bd0f00
·
verified ·
1 Parent(s): e2a5a9a

Update pk_workflow.py

Browse files
Files changed (1) hide show
  1. pk_workflow.py +0 -18
pk_workflow.py CHANGED
@@ -284,12 +284,6 @@ 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
- 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,12 +401,6 @@ def _dpmpp_2m_sde_step(scheduler, model_output, timestep, sample, eta: float = 1
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,12 +492,6 @@ def _dpmpp_3m_sde_step(scheduler, model_output, timestep, sample, eta: float = 1
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
 
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
 
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
 
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