diff --git a/nodes.py b/nodes.py index a48b3c5..4049b32 100644 --- a/nodes.py +++ b/nodes.py @@ -137,7 +137,7 @@ class KRestartSamplerCustom: "return_with_leftover_noise": (["disable", "enable"], ), "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(), ), - }, + } } RETURN_TYPES = ("LATENT","LATENT") @@ -172,7 +172,7 @@ class KRestartSamplerCustomNoise: "return_with_leftover_noise": (["disable", "enable"], ), "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(),), - "chunked_mode": (["disable", "enable"], ), + "chunked_mode": ("BOOLEAN", {"default": False}), }, "optional": { "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), @@ -187,7 +187,7 @@ class KRestartSamplerCustomNoise: def sample(self, model, add_noise, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler, custom_noise_opt=None, chunked_mode="disable"): force_full_denoise = return_with_leftover_noise != "enable" disable_noise = add_noise == "disable" - return restart_sampling(model, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, step_range=(start_at_step, end_at_step), force_full_denoise=force_full_denoise, output_only=False, custom_noise=custom_noise_opt, chunked_mode=chunked_mode=="enable") + return restart_sampling(model, noise_seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, segments, restart_scheduler, disable_noise=disable_noise, step_range=(start_at_step, end_at_step), force_full_denoise=force_full_denoise, output_only=False, custom_noise=custom_noise_opt.make_noise_sampler if custom_noise_opt else None, chunked_mode=chunked_mode) NODE_CLASS_MAPPINGS = { "KRestartSamplerSimple": KRestartSamplerSimple, diff --git a/restart_sampling.py b/restart_sampling.py index 2b29ed2..f6193cd 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -1,4 +1,5 @@ import ast +from collections import namedtuple import warnings import torch from tqdm.auto import trange @@ -173,24 +174,36 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega return (out, out_denoised) +class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart_sigmas"], defaults=[None, 0, 0., 0., None])): + __slots__ = () + + @torch.no_grad() + def execute(self, x, sample, get_noise_sampler): + x = sample(x, self.sigmas, -1) + if self.k < 1 or self.restart_sigmas is None: + return x + noise_sampler = get_noise_sampler(x, self.s_min, self.s_max) + for kidx in range(self.k): + x += noise_sampler(self.restart_sigmas[0], self.restart_sigmas[-1]) * (self.s_max ** 2 - self.s_min ** 2) ** 0.5 + x = sample(x, self.restart_sigmas, kidx) + return x + + class KSamplerRestartWrapper: - - ksampler = None - - def __init__(self, sampler, real_model, restart_scheduler, restart_segments, seed, custom_noise=None, chunked=True): + def __init__(self, sampler, real_model, restart_scheduler, restart_segments, seed, make_noise_sampler=None, chunked=True): self.ksampler = sampler self.real_model = real_model self.restart_scheduler = restart_scheduler self.restart_segments = restart_segments self.total_steps = 0 self.seed = seed - self.custom_noise = custom_noise + self.make_noise_sampler = make_noise_sampler self.chunked = chunked @torch.no_grad() - def build_plan(self, x, sigmas): + def build_plan(self, sigmas, device): segments = round_restart_segments(sigmas, self.restart_segments) - self.total_steps = len(sigmas) - 1 + calc_restart_steps(segments) + total_steps = len(sigmas) - 1 + calc_restart_steps(segments) plan = [] range_start = -1 for i in range(len(sigmas) - 1): @@ -202,20 +215,57 @@ class KSamplerRestartWrapper: continue s_max, k, n_restart = seg['t_max'], seg['k'], seg['n'] seg_sigmas = calc_sigmas(self.restart_scheduler, n_restart, s_min, - s_max, self.real_model, device=x.device) - plan.append((sigmas[range_start:i+2], k, s_min, s_max, seg_sigmas[:-1])) + s_max, self.real_model, device=device) + plan.append(PlanItem(sigmas[range_start:i+2], k, s_min, s_max, seg_sigmas[:-1])) range_start = -1 if range_start != -1: - plan.append((sigmas[range_start:], 0, 0, 0, None)) - return plan + plan.append(PlanItem(sigmas[range_start:])) + return plan, total_steps + + def explain_plan(self, plan, total_steps): + step = 0 + last_kidx = -1 + def do_sample(x, sigs, kidx=-1): + nonlocal step, last_kidx + rlabel = f"R{kidx+1:>3}" if kidx > last_kidx else " " + last_kidx = kidx + if not self.chunked: + for i in range(len(sigs)-1): + step += 1 + print(f"[{rlabel}] Step {step:>3}: {sigs[i:i+2]}") + return x + chunk_size = len(sigs) - 2 + step += 1 + print(f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {sigs}") + step += chunk_size + return x + + def get_noise_sampler(*_args): + return lambda *_args: 0.0 + + for pi in plan: + pi.execute(0.0, do_sample, get_noise_sampler) + @torch.no_grad() def ksampler_restart_wrapper(self, model, x, sigmas, *args, extra_args=None, callback=None, disable=None, **kwargs): ksampler = self.ksampler - def noise_sampler(_s, _sn): - return torch.randn_like(x) - plan = self.build_plan(x, sigmas) step = 0 + seed = self.seed + plan, self.total_steps = self.build_plan(sigmas, x.device) + + self.explain_plan(plan, self.total_steps) + + def noise_sampler(*_args): + return torch.randn_like(x) + + def get_noise_sampler(x, s_min, s_max): + nonlocal seed + if not self.make_noise_sampler: + return noise_sampler + result = self.make_noise_sampler(x, s_min, s_max, seed) + seed += 1 + return result with trange(self.total_steps, disable=disable) as pbar: def callback_wrapper(x): @@ -226,25 +276,21 @@ class KSamplerRestartWrapper: if callback is not None: callback(x) - def do_sample(x, sigs): - if isinstance(sigs, (list,tuple)): + def sampler_function(x, sigs): + return ksampler.sampler_function( + model, x, sigs, *args, extra_args=extra_args, callback=callback_wrapper, disable=True, + **kwargs) + + def do_sample(x, sigs, kidx=-1): + if isinstance(sigs, (list, tuple)): sigs = torch.tensor(sigs, device=x.device) if self.chunked or len(sigs) < 3: - return ksampler.sampler_function( - model, x, sigs, *args, extra_args=extra_args, callback=callback_wrapper, disable=True, - **kwargs) + return sampler_function(x, sigs) for i in range(len(sigs)-1): - x = do_sample(x, (sigs[i], sigs[i+1])) # This only ever recurses once. + x = sampler_function(x, sigs[i:i+2]) return x - for chunk_sigmas, k, s_min, s_max, restart_sigmas in plan: - x = do_sample(x, chunk_sigmas) - if restart_sigmas is None: - continue - if self.custom_noise is not None: - noise_sampler = self.custom_noise.make_noise_sampler(x, s_min, s_max, self.seed) - for _ in range(k): - x += noise_sampler(restart_sigmas[0], restart_sigmas[-1]) * (s_max ** 2 - s_min ** 2) ** 0.5 - x = do_sample(x, restart_sigmas) + for pi in plan: + x = pi.execute(x, do_sample, get_noise_sampler) return x