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