Spaces:
Running on Zero
Running on Zero
Update pk_workflow.py
Browse files- pk_workflow.py +13 -2
pk_workflow.py
CHANGED
|
@@ -354,7 +354,12 @@ def _dpmpp_2m_sde_step(scheduler, model_output, timestep, sample, eta: float = 1
|
|
| 354 |
denoised = sample + sigma_from_timestep * model_output
|
| 355 |
|
| 356 |
compute_dtype = torch.float32 if sample.dtype in (torch.float16, torch.bfloat16) else sample.dtype
|
| 357 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 358 |
sigma, sigma_next = sigmas[i], sigmas[i + 1]
|
| 359 |
x = sample.to(dtype=compute_dtype)
|
| 360 |
denoised = denoised.to(dtype=compute_dtype)
|
|
@@ -418,7 +423,12 @@ def _dpmpp_3m_sde_step(scheduler, model_output, timestep, sample, eta: float = 1
|
|
| 418 |
denoised = sample + sigma_from_timestep * model_output
|
| 419 |
|
| 420 |
compute_dtype = torch.float32 if sample.dtype in (torch.float16, torch.bfloat16) else sample.dtype
|
| 421 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 422 |
sigma, sigma_next = sigmas[i], sigmas[i + 1]
|
| 423 |
x = sample.to(dtype=compute_dtype)
|
| 424 |
denoised = denoised.to(dtype=compute_dtype)
|
|
@@ -552,6 +562,7 @@ class use_schedule:
|
|
| 552 |
scheduler._dpmpp_sde_h_last = None
|
| 553 |
scheduler._dpmpp_sde_h_last_2 = None
|
| 554 |
scheduler._dpmpp_sde_noise_sampler = None
|
|
|
|
| 555 |
scheduler._dpmpp_sde_seed = self.seed + offset
|
| 556 |
|
| 557 |
def stepped(model_output, timestep, sample, return_dict=True, _s=scheduler, _f=step_fn, **_kwargs):
|
|
|
|
| 354 |
denoised = sample + sigma_from_timestep * model_output
|
| 355 |
|
| 356 |
compute_dtype = torch.float32 if sample.dtype in (torch.float16, torch.bfloat16) else sample.dtype
|
| 357 |
+
# Cached once and reused every call — `torchsde.BrownianTree` caches its internal tree keyed to the exact
|
| 358 |
+
# float value it was first queried with, and re-deriving "the same" sigma via a fresh `.to()` cast on a
|
| 359 |
+
# later call can land a few ULPs away from what the tree remembers, which it treats as an ordering error.
|
| 360 |
+
if scheduler._dpmpp_sde_sigmas is None:
|
| 361 |
+
scheduler._dpmpp_sde_sigmas = scheduler.sigmas.to(device=sample.device, dtype=compute_dtype)
|
| 362 |
+
sigmas = scheduler._dpmpp_sde_sigmas
|
| 363 |
sigma, sigma_next = sigmas[i], sigmas[i + 1]
|
| 364 |
x = sample.to(dtype=compute_dtype)
|
| 365 |
denoised = denoised.to(dtype=compute_dtype)
|
|
|
|
| 423 |
denoised = sample + sigma_from_timestep * model_output
|
| 424 |
|
| 425 |
compute_dtype = torch.float32 if sample.dtype in (torch.float16, torch.bfloat16) else sample.dtype
|
| 426 |
+
# Cached once and reused every call — `torchsde.BrownianTree` caches its internal tree keyed to the exact
|
| 427 |
+
# float value it was first queried with, and re-deriving "the same" sigma via a fresh `.to()` cast on a
|
| 428 |
+
# later call can land a few ULPs away from what the tree remembers, which it treats as an ordering error.
|
| 429 |
+
if scheduler._dpmpp_sde_sigmas is None:
|
| 430 |
+
scheduler._dpmpp_sde_sigmas = scheduler.sigmas.to(device=sample.device, dtype=compute_dtype)
|
| 431 |
+
sigmas = scheduler._dpmpp_sde_sigmas
|
| 432 |
sigma, sigma_next = sigmas[i], sigmas[i + 1]
|
| 433 |
x = sample.to(dtype=compute_dtype)
|
| 434 |
denoised = denoised.to(dtype=compute_dtype)
|
|
|
|
| 562 |
scheduler._dpmpp_sde_h_last = None
|
| 563 |
scheduler._dpmpp_sde_h_last_2 = None
|
| 564 |
scheduler._dpmpp_sde_noise_sampler = None
|
| 565 |
+
scheduler._dpmpp_sde_sigmas = None
|
| 566 |
scheduler._dpmpp_sde_seed = self.seed + offset
|
| 567 |
|
| 568 |
def stepped(model_output, timestep, sample, return_dict=True, _s=scheduler, _f=step_fn, **_kwargs):
|