Spaces:
Running on Zero
Running on Zero
Update pk_workflow.py
Browse files- 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
|