fix res_multistep and its ancestral.

This commit is contained in:
Balladie 2025-08-20 17:16:17 +09:00
parent 5a8f502db5
commit 4ce6a1dd99

View File

@ -1342,12 +1342,16 @@ def res_multistep(model, x, sigmas, extra_args=None, callback=None, disable=None
seed = extra_args.get("seed", None) seed = extra_args.get("seed", None)
noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler
s_in = x.new_ones([x.shape[0]]) s_in = x.new_ones([x.shape[0]])
sigma_fn = lambda t: t.neg().exp()
t_fn = lambda sigma: sigma.log().neg() model_sampling = model.inner_model.model_patcher.get_model_object('model_sampling')
sigma_fn = partial(half_log_snr_to_sigma, model_sampling=model_sampling)
lambda_fn = partial(sigma_to_half_log_snr, model_sampling=model_sampling)
sigmas = offset_first_sigma_for_snr(sigmas, model_sampling)
phi1_fn = lambda t: torch.expm1(t) / t phi1_fn = lambda t: torch.expm1(t) / t
phi2_fn = lambda t: (phi1_fn(t) - 1.0) / t phi2_fn = lambda t: (phi1_fn(t) - 1.0) / t
old_sigma_down = None old_sigma_next = None
old_denoised = None old_denoised = None
uncond_denoised = None uncond_denoised = None
def post_cfg_function(args): def post_cfg_function(args):
@ -1361,43 +1365,46 @@ def res_multistep(model, x, sigmas, extra_args=None, callback=None, disable=None
for i in trange(len(sigmas) - 1, disable=disable): for i in trange(len(sigmas) - 1, disable=disable):
denoised = model(x, sigmas[i] * s_in, **extra_args) denoised = model(x, sigmas[i] * s_in, **extra_args)
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) # sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
if callback is not None: if callback is not None:
callback({"x": x, "i": i, "sigma": sigmas[i], "sigma_hat": sigmas[i], "denoised": denoised}) callback({"x": x, "i": i, "sigma": sigmas[i], "sigma_hat": sigmas[i], "denoised": denoised})
if sigma_down == 0 or old_denoised is None: if sigmas[i + 1] == 0 or old_denoised is None:
# Euler method # Euler method
if cfg_pp: if cfg_pp:
d = to_d(x, sigmas[i], uncond_denoised) d = to_d(x, sigmas[i], uncond_denoised)
x = denoised + d * sigma_down x = denoised + d * sigmas[i + 1]
else: else:
d = to_d(x, sigmas[i], denoised) d = to_d(x, sigmas[i], denoised)
dt = sigma_down - sigmas[i] dt = sigmas[i + 1] - sigmas[i]
x = x + d * dt x = x + d * dt
else: else:
# Second order multistep method in https://arxiv.org/pdf/2308.02157 # Second order multistep method in https://arxiv.org/pdf/2308.02157
t, t_old, t_next, t_prev = t_fn(sigmas[i]), t_fn(old_sigma_down), t_fn(sigma_down), t_fn(sigmas[i - 1]) t, t_old, t_next, t_prev = lambda_fn(sigmas[i]), lambda_fn(old_sigma_next), lambda_fn(sigmas[i + 1]), lambda_fn(sigmas[i - 1])
h = t_next - t h = t_next - t
h_eta = h * (eta + 1)
c2 = (t_prev - t_old) / h c2 = (t_prev - t_old) / h
phi1_val, phi2_val = phi1_fn(-h), phi2_fn(-h) alpha_next = sigmas[i + 1] * t_next.exp()
phi1_val, phi2_val = phi1_fn(-h_eta), phi2_fn(-h_eta)
b1 = torch.nan_to_num(phi1_val - phi2_val / c2, nan=0.0) b1 = torch.nan_to_num(phi1_val - phi2_val / c2, nan=0.0)
b2 = torch.nan_to_num(phi2_val / c2, nan=0.0) b2 = torch.nan_to_num(phi2_val / c2, nan=0.0)
if cfg_pp: if cfg_pp:
x = x + (denoised - uncond_denoised) x = x + (denoised - uncond_denoised)
x = sigma_fn(h) * x + h * (b1 * uncond_denoised + b2 * old_denoised) x = sigmas[i + 1] / sigmas[i] * (-h * eta).exp() * x + alpha_next * h_eta * (b1 * uncond_denoised + b2 * old_denoised)
else: else:
x = sigma_fn(h) * x + h * (b1 * denoised + b2 * old_denoised) x = sigmas[i + 1] / sigmas[i] * (-h * eta).exp() * x + alpha_next * h_eta * (b1 * denoised + b2 * old_denoised)
# Noise addition # Noise addition
if sigmas[i + 1] > 0: sigma_up = sigmas[i + 1] * (-2 * h * eta).expm1().neg().sqrt()
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
if cfg_pp: if cfg_pp:
old_denoised = uncond_denoised old_denoised = uncond_denoised
else: else:
old_denoised = denoised old_denoised = denoised
old_sigma_down = sigma_down old_sigma_next = sigmas[i + 1]
return x return x
@torch.no_grad() @torch.no_grad()