dagloop5 commited on
Commit
4d7acb7
·
verified ·
1 Parent(s): 17335fd

Update pk_workflow.py

Browse files
Files changed (1) hide show
  1. pk_workflow.py +88 -0
pk_workflow.py CHANGED
@@ -217,6 +217,83 @@ def _euler_ancestral_step(scheduler, generator, model_output, timestep, sample,
217
  scheduler._step_index += 1
218
  return prev_sample
219
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
220
  class use_schedule:
221
  """Set each scheduler's shift for one request, and — for anything but `native` — force its sigma grid onto
222
  one of `SCHEDULE_SIGMA_FUNCS`'s named schedules.
@@ -269,6 +346,17 @@ class use_schedule:
269
  def stepped(model_output, timestep, sample, return_dict=True, _s=scheduler, _g=generator, **_kwargs):
270
  return (_euler_ancestral_step(_s, _g, model_output, timestep, sample),)
271
 
 
 
 
 
 
 
 
 
 
 
 
272
  scheduler.step = stepped
273
  return self
274
 
 
217
  scheduler._step_index += 1
218
  return prev_sample
219
 
220
+ def _er_sde_step(scheduler, generator, model_output, timestep, sample, s_noise: float = 1.0, max_stage: int = 3):
221
+ """Ports k-diffusion's `sample_er_sde` (VP ER-SDE-Solver-3, arXiv:2309.06169) onto one
222
+ `MiniMaxH3Scheduler.step()` call. Single model evaluation per step — second/third-order accuracy comes from
223
+ the previous one or two steps' denoised estimates, not an extra evaluation this step — so it carries history
224
+ on the scheduler instance across calls, reset each request by `use_schedule` alongside `_step_index`.
225
+ """
226
+ if scheduler._step_index is None:
227
+ scheduler._step_index = scheduler.index_for_timestep(timestep) if scheduler._begin_index is None else scheduler._begin_index
228
+ i = scheduler._step_index
229
+
230
+ if not isinstance(timestep, torch.Tensor):
231
+ timestep = torch.tensor(timestep, dtype=sample.dtype)
232
+ sigma_from_timestep = 1 - timestep.to(device=sample.device, dtype=sample.dtype)
233
+ while sigma_from_timestep.ndim < sample.ndim:
234
+ sigma_from_timestep = sigma_from_timestep.unsqueeze(-1)
235
+ denoised = sample + sigma_from_timestep * model_output
236
+
237
+ compute_dtype = torch.float32 if sample.dtype in (torch.float16, torch.bfloat16) else sample.dtype
238
+ sigmas = scheduler.sigmas.to(device=sample.device, dtype=compute_dtype)
239
+ sigma, sigma_next = sigmas[i], sigmas[i + 1]
240
+ x = sample.to(dtype=compute_dtype)
241
+ denoised = denoised.to(dtype=compute_dtype)
242
+
243
+ if i == 0 and float(sigma) >= 1.0:
244
+ # `1 - sigma` sits in a denominator below; MiniMax-H3's first sigma is exactly 1.0, so nudge it a hair
245
+ # under 1.0 for this sampler's math only, matching ComfyUI's `offset_first_sigma_for_snr`. Does not
246
+ # touch `sigma_from_timestep` above — the model was still conditioned on the real timestep.
247
+ base = torch.tensor(1.0 - 1e-4, dtype=compute_dtype, device=sample.device)
248
+ shift = float(scheduler.shift)
249
+ sigma = shift * base / (1 + (shift - 1) * base)
250
+
251
+ def er_lambda(s):
252
+ return s / (1 - s)
253
+
254
+ def noise_scaler(v):
255
+ return v * (v**0.3).exp() + v * 10.0
256
+
257
+ if sigma_next == 0:
258
+ prev_sample = denoised
259
+ else:
260
+ er_lambda_s, er_lambda_t = er_lambda(sigma), er_lambda(sigma_next)
261
+ alpha_s, alpha_t = 1 - sigma, 1 - sigma_next
262
+ r_alpha = alpha_t / alpha_s
263
+ r = noise_scaler(er_lambda_t) / noise_scaler(er_lambda_s)
264
+
265
+ prev_sample = r_alpha * r * x + alpha_t * (1 - r) * denoised
266
+
267
+ stage_used = min(max_stage, i + 1)
268
+ if stage_used >= 2:
269
+ num_points = 200
270
+ dt = er_lambda_t - er_lambda_s
271
+ step_size = -dt / num_points
272
+ positions = er_lambda_t + torch.arange(num_points, device=x.device, dtype=compute_dtype) * step_size
273
+ scaled = noise_scaler(positions)
274
+
275
+ s_term = torch.sum(1 / scaled) * step_size
276
+ er_lambda_prev = er_lambda(sigmas[i - 1])
277
+ denoised_d = (denoised - scheduler._er_sde_old_denoised) / (er_lambda_s - er_lambda_prev)
278
+ prev_sample = prev_sample + alpha_t * (dt + s_term * noise_scaler(er_lambda_t)) * denoised_d
279
+
280
+ if stage_used >= 3:
281
+ s_u_term = torch.sum((positions - er_lambda_s) / scaled) * step_size
282
+ er_lambda_prev2 = er_lambda(sigmas[i - 2])
283
+ denoised_u = (denoised_d - scheduler._er_sde_old_denoised_d) / ((er_lambda_s - er_lambda_prev2) / 2)
284
+ prev_sample = prev_sample + alpha_t * ((dt**2) / 2 + s_u_term * noise_scaler(er_lambda_t)) * denoised_u
285
+ scheduler._er_sde_old_denoised_d = denoised_d
286
+
287
+ if s_noise > 0:
288
+ noise = torch.randn(x.shape, dtype=x.dtype, device="cpu", generator=generator).to(x.device)
289
+ spread = (er_lambda_t**2 - er_lambda_s**2 * r**2).clamp_min(0).sqrt()
290
+ prev_sample = prev_sample + alpha_t * noise * s_noise * spread
291
+
292
+ scheduler._er_sde_old_denoised = denoised
293
+ prev_sample = prev_sample.to(dtype=sample.dtype)
294
+ scheduler._step_index += 1
295
+ return prev_sample
296
+
297
  class use_schedule:
298
  """Set each scheduler's shift for one request, and — for anything but `native` — force its sigma grid onto
299
  one of `SCHEDULE_SIGMA_FUNCS`'s named schedules.
 
346
  def stepped(model_output, timestep, sample, return_dict=True, _s=scheduler, _g=generator, **_kwargs):
347
  return (_euler_ancestral_step(_s, _g, model_output, timestep, sample),)
348
 
349
+ scheduler.step = stepped
350
+ elif self.sampler_name == "er_sde":
351
+ for offset, attr_name in enumerate(self.attr_names):
352
+ scheduler = getattr(self.pipe, attr_name)
353
+ scheduler._er_sde_old_denoised = None
354
+ scheduler._er_sde_old_denoised_d = None
355
+ generator = torch.Generator(device="cpu").manual_seed(self.seed + offset)
356
+
357
+ def stepped(model_output, timestep, sample, return_dict=True, _s=scheduler, _g=generator, **_kwargs):
358
+ return (_er_sde_step(_s, _g, model_output, timestep, sample),)
359
+
360
  scheduler.step = stepped
361
  return self
362