dagloop5 commited on
Commit
e70f255
·
verified ·
1 Parent(s): 5d1d6c8

Update pk_workflow.py

Browse files
Files changed (1) hide show
  1. 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
- sigmas = scheduler.sigmas.to(device=sample.device, dtype=compute_dtype)
 
 
 
 
 
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
- sigmas = scheduler.sigmas.to(device=sample.device, dtype=compute_dtype)
 
 
 
 
 
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):