From 22f4ee04e5fa97902c862b6ff359025dbb4ceb55 Mon Sep 17 00:00:00 2001 From: blepping Date: Sat, 16 Mar 2024 07:21:41 -0600 Subject: [PATCH 01/10] Hack to allow setting custom noise --- nodes.py | 9 ++++++--- restart_sampling.py | 13 +++++++++---- 2 files changed, 15 insertions(+), 7 deletions(-) diff --git a/nodes.py b/nodes.py index 269c86f..785f8c0 100644 --- a/nodes.py +++ b/nodes.py @@ -137,7 +137,10 @@ class KRestartSamplerCustom: "return_with_leftover_noise": (["disable", "enable"], ), "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(), ), - } + }, + "optional": { + "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), + }, } RETURN_TYPES = ("LATENT","LATENT") @@ -145,10 +148,10 @@ class KRestartSamplerCustom: FUNCTION = "sample" CATEGORY = "sampling" - 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): + 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): 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) + 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) NODE_CLASS_MAPPINGS = { diff --git a/restart_sampling.py b/restart_sampling.py index 771730c..fe71bc2 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -87,7 +87,7 @@ def calc_restart_steps(restart_segments): return restart_steps -def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True): +def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None): if isinstance(sampler, str): sampler = sampler_object(sampler) @@ -118,7 +118,7 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega sigmas = sigmas[-(steps + 1):] total_steps = [0] # Updated in the wrapper. - sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, total_steps) + sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise) latent = latent_image latent_image = latent["samples"] @@ -178,16 +178,19 @@ class KSamplerRestartWrapper: ksampler = None - def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps): + def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise): self.ksampler = sampler self.real_model = real_model self.restart_scheduler = restart_scheduler self.restart_segments = restart_segments self.total_steps = total_steps + self.seed = seed + self.custom_noise = custom_noise @torch.no_grad() def ksampler_restart_wrapper(self, model, x, sigmas, *args, extra_args=None, callback=None, disable=None, **kwargs): ksampler = self.ksampler + noise_sampler = lambda _s, _sn: torch.randn_like(x) segments = round_restart_segments(sigmas, self.restart_segments) self.total_steps[0] = len(sigmas) - 1 + calc_restart_steps(segments) step = 0 @@ -211,8 +214,10 @@ class KSamplerRestartWrapper: 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) + 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 += torch.randn_like(x) * (s_max ** 2 - s_min ** 2) ** 0.5 + x += noise_sampler(None, None) * (s_max ** 2 - s_min ** 2) ** 0.5 for j in range(n_restart - 1): x = ksampler.sampler_function(model, x, torch.tensor( [seg_sigmas[j], seg_sigmas[j + 1]], device=x.device), *args, extra_args=extra_args, From 33ba61bb78db402c673802a558f28023681d3a2d Mon Sep 17 00:00:00 2001 From: blepping Date: Fri, 22 Mar 2024 11:47:08 -0600 Subject: [PATCH 02/10] Make custom restart sampler with custom noise a separate node --- nodes.py | 37 ++++++++++++++++++++++++++++++++++++- restart_sampling.py | 9 +++++---- 2 files changed, 41 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index 785f8c0..116b24b 100644 --- a/nodes.py +++ b/nodes.py @@ -138,6 +138,41 @@ class KRestartSamplerCustom: "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(), ), }, + } + + RETURN_TYPES = ("LATENT","LATENT") + RETURN_NAMES = ("output", "denoised_output") + FUNCTION = "sample" + CATEGORY = "sampling" + + 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): + 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) + + +class KRestartSamplerCustomNoise: + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "add_noise": (["enable", "disable"], ), + "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), + "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}), + "sampler": ("SAMPLER", ), + "scheduler": (tuple(SCHEDULER_MAPPING.keys()), ), + "positive": ("CONDITIONING", ), + "negative": ("CONDITIONING", ), + "latent_image": ("LATENT", ), + "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), + "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}), + "return_with_leftover_noise": (["disable", "enable"], ), + "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), + "restart_scheduler": (get_supported_restart_schedulers(),), + }, "optional": { "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), }, @@ -153,12 +188,12 @@ class KRestartSamplerCustom: 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) - NODE_CLASS_MAPPINGS = { "KRestartSamplerSimple": KRestartSamplerSimple, "KRestartSampler": KRestartSampler, "KRestartSamplerAdv": KRestartSamplerAdv, "KRestartSamplerCustom": KRestartSamplerCustom, + "KRestartSamplerCustomNoise": KRestartSamplerCustomNoise, } NODE_DISPLAY_NAME_MAPPINGS = { diff --git a/restart_sampling.py b/restart_sampling.py index fe71bc2..739168a 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -87,7 +87,7 @@ def calc_restart_steps(restart_segments): return restart_steps -def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None): +def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, noise_multiplier=1.0): if isinstance(sampler, str): sampler = sampler_object(sampler) @@ -178,7 +178,7 @@ class KSamplerRestartWrapper: ksampler = None - def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise): + def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise=None): self.ksampler = sampler self.real_model = real_model self.restart_scheduler = restart_scheduler @@ -190,7 +190,8 @@ class KSamplerRestartWrapper: @torch.no_grad() def ksampler_restart_wrapper(self, model, x, sigmas, *args, extra_args=None, callback=None, disable=None, **kwargs): ksampler = self.ksampler - noise_sampler = lambda _s, _sn: torch.randn_like(x) + def noise_sampler(_s, _sn): + return torch.randn_like(x) segments = round_restart_segments(sigmas, self.restart_segments) self.total_steps[0] = len(sigmas) - 1 + calc_restart_steps(segments) step = 0 @@ -217,7 +218,7 @@ class KSamplerRestartWrapper: 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(None, None) * (s_max ** 2 - s_min ** 2) ** 0.5 + x += noise_sampler(seg_sigmas[0],seg_sigmas[-1]) * (s_max ** 2 - s_min ** 2) ** 0.5 for j in range(n_restart - 1): x = ksampler.sampler_function(model, x, torch.tensor( [seg_sigmas[j], seg_sigmas[j + 1]], device=x.device), *args, extra_args=extra_args, From b172908ac784576a9371e9085e9895935407efd1 Mon Sep 17 00:00:00 2001 From: blepping Date: Fri, 22 Mar 2024 14:45:04 -0600 Subject: [PATCH 03/10] Implement chunked restart sampling --- nodes.py | 5 +-- restart_sampling.py | 79 +++++++++++++++++++++++++++++---------------- 2 files changed, 54 insertions(+), 30 deletions(-) diff --git a/nodes.py b/nodes.py index 116b24b..a48b3c5 100644 --- a/nodes.py +++ b/nodes.py @@ -172,6 +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"], ), }, "optional": { "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), @@ -183,10 +184,10 @@ class KRestartSamplerCustomNoise: FUNCTION = "sample" CATEGORY = "sampling" - 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): + 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) + 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") NODE_CLASS_MAPPINGS = { "KRestartSamplerSimple": KRestartSamplerSimple, diff --git a/restart_sampling.py b/restart_sampling.py index 739168a..bfc482f 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -87,7 +87,7 @@ def calc_restart_steps(restart_segments): return restart_steps -def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, noise_multiplier=1.0): +def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, noise_multiplier=1.0, chunked_mode=False): if isinstance(sampler, str): sampler = sampler_object(sampler) @@ -118,7 +118,7 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega sigmas = sigmas[-(steps + 1):] total_steps = [0] # Updated in the wrapper. - sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise) + sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise, chunked=chunked_mode) latent = latent_image latent_image = latent["samples"] @@ -178,7 +178,7 @@ class KSamplerRestartWrapper: ksampler = None - def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise=None): + def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise=None, chunked=True): self.ksampler = sampler self.real_model = real_model self.restart_scheduler = restart_scheduler @@ -186,43 +186,66 @@ class KSamplerRestartWrapper: self.total_steps = total_steps self.seed = seed self.custom_noise = custom_noise + self.chunked = chunked + + @torch.no_grad() + def build_plan(self, x, sigmas): + segments = round_restart_segments(sigmas, self.restart_segments) + self.total_steps[0] = len(sigmas) - 1 + calc_restart_steps(segments) + plan = [] + range_start = -1 + for i in range(len(sigmas) - 1): + if range_start == -1: + range_start = i + s_min = sigmas[i + 1].item() + seg = segments.get(s_min) + if seg is None: + 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])) + range_start = -1 + if range_start != -1: + plan.append((sigmas[range_start:], 0, 0, 0, None)) + return plan @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) - segments = round_restart_segments(sigmas, self.restart_segments) - self.total_steps[0] = len(sigmas) - 1 + calc_restart_steps(segments) + plan = self.build_plan(x, sigmas) step = 0 - def callback_wrapper(x): - x["i"] = step - if callback is not None: - callback(x) with trange(self.total_steps[0], disable=disable) as pbar: - for i in range(len(sigmas) - 1): - x = ksampler.sampler_function( - model, x, torch.tensor([sigmas[i], sigmas[i + 1]], - device=x.device), *args, extra_args=extra_args, callback=callback_wrapper, disable=True, - **kwargs) - pbar.update(1) + def callback_wrapper(x): + nonlocal step step += 1 - s_min = sigmas[i + 1].item() - seg = segments.get(s_min) - if seg is None: + pbar.update(1) + x["i"] = step + if callback is not None: + callback(x) + + def do_sample(x, sigs): + 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) + for i in range(len(sigs)-1): + x = do_sample(x, (sigs[i], sigs[i+1])) # This only ever recurses once. + 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 - 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) 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(seg_sigmas[0],seg_sigmas[-1]) * (s_max ** 2 - s_min ** 2) ** 0.5 - for j in range(n_restart - 1): - x = ksampler.sampler_function(model, x, torch.tensor( - [seg_sigmas[j], seg_sigmas[j + 1]], device=x.device), *args, extra_args=extra_args, - callback=callback_wrapper, disable=True, **kwargs) - pbar.update(1) - step += 1 + x += noise_sampler(restart_sigmas[0], restart_sigmas[-1]) * (s_max ** 2 - s_min ** 2) ** 0.5 + x = do_sample(x, restart_sigmas) + return x From 12c0ca1044ea2846c3904f5d65b59d741b7c4aad Mon Sep 17 00:00:00 2001 From: blepping Date: Fri, 22 Mar 2024 14:51:20 -0600 Subject: [PATCH 04/10] Simplify total steps logic in wrapper --- restart_sampling.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/restart_sampling.py b/restart_sampling.py index bfc482f..2b29ed2 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -117,8 +117,7 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega elif effective_steps != steps: sigmas = sigmas[-(steps + 1):] - total_steps = [0] # Updated in the wrapper. - sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise, chunked=chunked_mode) + sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, seed, custom_noise, chunked=chunked_mode) latent = latent_image latent_image = latent["samples"] @@ -148,7 +147,7 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega pbar_update_absolute = ProgressBar.update_absolute def pbar_update_absolute_wrapper(self, value, total=None, preview=None): - pbar_update_absolute(self, value, total_steps[0], preview) + pbar_update_absolute(self, value, sampler_wrapper.total_steps, preview) ProgressBar.update_absolute = pbar_update_absolute_wrapper @@ -178,12 +177,12 @@ class KSamplerRestartWrapper: ksampler = None - def __init__(self, sampler, real_model, restart_scheduler, restart_segments, total_steps, seed, custom_noise=None, chunked=True): + def __init__(self, sampler, real_model, restart_scheduler, restart_segments, seed, custom_noise=None, chunked=True): self.ksampler = sampler self.real_model = real_model self.restart_scheduler = restart_scheduler self.restart_segments = restart_segments - self.total_steps = total_steps + self.total_steps = 0 self.seed = seed self.custom_noise = custom_noise self.chunked = chunked @@ -191,7 +190,7 @@ class KSamplerRestartWrapper: @torch.no_grad() def build_plan(self, x, sigmas): segments = round_restart_segments(sigmas, self.restart_segments) - self.total_steps[0] = len(sigmas) - 1 + calc_restart_steps(segments) + self.total_steps = len(sigmas) - 1 + calc_restart_steps(segments) plan = [] range_start = -1 for i in range(len(sigmas) - 1): @@ -218,7 +217,7 @@ class KSamplerRestartWrapper: plan = self.build_plan(x, sigmas) step = 0 - with trange(self.total_steps[0], disable=disable) as pbar: + with trange(self.total_steps, disable=disable) as pbar: def callback_wrapper(x): nonlocal step step += 1 From 67d4b62235a0e6187c06c1fb1bd3db7842d8ecb8 Mon Sep 17 00:00:00 2001 From: blepping Date: Mon, 25 Mar 2024 05:45:48 -0600 Subject: [PATCH 05/10] Refactor plan handling --- nodes.py | 6 +-- restart_sampling.py | 104 ++++++++++++++++++++++++++++++++------------ 2 files changed, 78 insertions(+), 32 deletions(-) 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 From bbbddbd7cbdd34f72995c99edeb4beffafe12510 Mon Sep 17 00:00:00 2001 From: blepping Date: Mon, 25 Mar 2024 06:00:09 -0600 Subject: [PATCH 06/10] Remove unused noise_multiplier param --- restart_sampling.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/restart_sampling.py b/restart_sampling.py index f6193cd..5605079 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -88,7 +88,7 @@ def calc_restart_steps(restart_segments): return restart_steps -def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, noise_multiplier=1.0, chunked_mode=False): +def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, chunked_mode=False): if isinstance(sampler, str): sampler = sampler_object(sampler) From e2e25e7a2870e0cacabc01e28106df2d06acea42 Mon Sep 17 00:00:00 2001 From: blepping Date: Wed, 27 Mar 2024 05:45:06 -0600 Subject: [PATCH 07/10] Add some comments describing what's going on in plan generation and execution Allow specifying segments "default" to use the default segments Allow specifying segments "a1111" to calculate segments like A1111 Allow setting environment variable COMFYUI_VERBOSE_RESTART_SAMPLING=1 to get some debug info Add chunked_mode to samplers (except for simple which will use the default of True) Documentation updates --- README.md | 18 +++++++++++ nodes.py | 63 +++++++----------------------------- restart_sampling.py | 79 +++++++++++++++++++++++++++++++++++++++++---- 3 files changed, 103 insertions(+), 57 deletions(-) diff --git a/README.md b/README.md index 905ba0c..9a476fc 100644 --- a/README.md +++ b/README.md @@ -16,6 +16,9 @@ git clone https://github.com/ssitu/ComfyUI_restart_sampling The Restart sampler nodes can be found in the node menu under `sampling`. +If you set the environment variable `COMFYUI_VERBOSE_RESTART_SAMPLING` to `1`, restart sampling will dump +information about the steps it's going to run to the console. + ### Nodes |Node|Image|Description| @@ -37,6 +40,21 @@ Both $t_{\textrm{min}}$ and $t_{\textrm{max}}$ within a segment definition may b You may freely mix the different formats. For example, `[2, 2, -500, "10%"], [3, 2, 5.3, -3]` would be a valid sequence. Note: Random numbers used for example only, not recommended. +**Special segment values**: + +* Enter `default` to use the default segment list. +* Enter `a1111` to emulate A1111 WebUI's segment calculation behavior. +For full emulation, enabled chunked mode, set both schedulers to `karras` and the sampler to `heun`. + +### Chunked Mode + +When chunked mode is enabled, the sampler is called with as many steps as possible up to the next segment. When disabled, the sampler +is only called with a single step at a time. Some samplers such as SDE samplers, momentum samplers, second order samplers +like dpmpp_2m use state from previous steps - when called step-by-step, this state is lost. Using chunked mode may make those +samplers more accurate. + +*Note*: Using SDE or momentum samplers with restart is likely not an improvement over normal sampling. + ## Visual Example Consider the default segments of `[3,2,0.06,0.30],[3,1,0.30,0.59]`. diff --git a/nodes.py b/nodes.py index 4049b32..b5914d4 100644 --- a/nodes.py +++ b/nodes.py @@ -1,5 +1,5 @@ import comfy -from .restart_sampling import restart_sampling, SCHEDULER_MAPPING +from .restart_sampling import restart_sampling, SCHEDULER_MAPPING, DEFAULT_SEGMENTS def get_supported_samplers(): @@ -24,8 +24,6 @@ def get_supported_samplers(): def get_supported_restart_schedulers(): return list(SCHEDULER_MAPPING.keys()) -DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]" - class KRestartSamplerSimple: @classmethod @@ -42,7 +40,7 @@ class KRestartSamplerSimple: "negative": ("CONDITIONING", ), "latent_image": ("LATENT", ), "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), - "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), + "segments": ("STRING", {"default": "default", "multiline": False}), } } @@ -50,7 +48,7 @@ class KRestartSamplerSimple: FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments): + def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, chunked_mode=False): return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, scheduler, denoise=denoise) @@ -71,6 +69,7 @@ class KRestartSampler: "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(), ), + "chunked_mode": ("BOOLEAN", {"default": True}), } } @@ -78,8 +77,8 @@ class KRestartSampler: FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler): - return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise) + def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler, chunked_mode=False): + return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise, chunked_mode=chunked_mode) class KRestartSamplerAdv: @@ -103,6 +102,7 @@ class KRestartSamplerAdv: "return_with_leftover_noise": (["disable", "enable"], ), "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(), ), + "chunked_mode": ("BOOLEAN", {"default": True}), } } @@ -110,10 +110,10 @@ class KRestartSamplerAdv: FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler): + def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler, chunked_mode=False): force_full_denoise = return_with_leftover_noise != "enable" disable_noise = add_noise == "disable" - return restart_sampling(model, noise_seed, steps, cfg, sampler_name, 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) + return restart_sampling(model, noise_seed, steps, cfg, sampler_name, 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, chunked_mode=chunked_mode) class KRestartSamplerCustom: @@ -137,6 +137,7 @@ class KRestartSamplerCustom: "return_with_leftover_noise": (["disable", "enable"], ), "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), "restart_scheduler": (get_supported_restart_schedulers(), ), + "chunked_mode": ("BOOLEAN", {"default": True}), } } @@ -145,56 +146,16 @@ class KRestartSamplerCustom: FUNCTION = "sample" CATEGORY = "sampling" - 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): + 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, chunked_mode=False): 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) - - -class KRestartSamplerCustomNoise: - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "add_noise": (["enable", "disable"], ), - "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), - "cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}), - "sampler": ("SAMPLER", ), - "scheduler": (tuple(SCHEDULER_MAPPING.keys()), ), - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "latent_image": ("LATENT", ), - "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), - "end_at_step": ("INT", {"default": 10000, "min": 0, "max": 10000}), - "return_with_leftover_noise": (["disable", "enable"], ), - "segments": ("STRING", {"default": DEFAULT_SEGMENTS, "multiline": False}), - "restart_scheduler": (get_supported_restart_schedulers(),), - "chunked_mode": ("BOOLEAN", {"default": False}), - }, - "optional": { - "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), - }, - } - - RETURN_TYPES = ("LATENT","LATENT") - RETURN_NAMES = ("output", "denoised_output") - FUNCTION = "sample" - CATEGORY = "sampling" - - 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.make_noise_sampler if custom_noise_opt else None, chunked_mode=chunked_mode) + 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, chunked_mode=chunked_mode) NODE_CLASS_MAPPINGS = { "KRestartSamplerSimple": KRestartSamplerSimple, "KRestartSampler": KRestartSampler, "KRestartSamplerAdv": KRestartSamplerAdv, "KRestartSamplerCustom": KRestartSamplerCustom, - "KRestartSamplerCustomNoise": KRestartSamplerCustomNoise, } NODE_DISPLAY_NAME_MAPPINGS = { diff --git a/restart_sampling.py b/restart_sampling.py index 5605079..143f718 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -1,5 +1,6 @@ import ast from collections import namedtuple +import os import warnings import torch from tqdm.auto import trange @@ -10,6 +11,10 @@ from comfy.samplers import KSAMPLER, sampler_object from comfy.utils import ProgressBar from .restart_schedulers import SCHEDULER_MAPPING +VERBOSE = os.environ.get("COMFYUI_VERBOSE_RESTART_SAMPLING", "").strip() == "1" + +DEFAULT_SEGMENTS = "[3,2,0.06,0.30],[3,1,0.30,0.59]" + def add_restart_segment(restart_segments, n_restart, k, t_min, t_max): if restart_segments is None: @@ -34,7 +39,25 @@ def resolve_t_value(val, ms): raise ValueError("bad t_min or t_max value") -def prepare_restart_segments(restart_info, ms): +def prepare_restart_segments(restart_info, ms, sigmas): + restart_info = restart_info.strip().lower() + if restart_info == "": + # No restarts. + return [] + if restart_info == "default": + restart_info = DEFAULT_SEGMENTS + elif restart_info == "a1111": + # Emulate A1111 WebUI's restart sampler behavior. + steps = len(sigmas) - 1 + if steps < 20: + # Less than 20 steps - no restarts. + return [] + if steps < 36: + # Less than 36 steps - one restart with 9 steps. + restart_info = "[10,1,0.1,0.2]" + else: + # Otherwise two restarts with steps // 4 steps. + restart_info = f"[{(steps // 4) + 1}, 2, 0.1, 0.2]" try: restart_arrays = ast.literal_eval(f"[{restart_info}]") except SyntaxError as e: @@ -88,18 +111,15 @@ def calc_restart_steps(restart_segments): return restart_steps -def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, chunked_mode=False): +def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, chunked_mode=True): if isinstance(sampler, str): sampler = sampler_object(sampler) - comfy.model_management.load_models_gpu([model]) real_model = model while hasattr(real_model, "model"): real_model = real_model.model - restart_segments = prepare_restart_segments(restart_info, real_model.model_sampling) - effective_steps = steps if step_range is not None or denoise > 0.9999 else int(steps / denoise) sigmas = calc_sigmas(scheduler, effective_steps, float(real_model.model_sampling.sigma_min), float(real_model.model_sampling.sigma_max), @@ -118,6 +138,8 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega elif effective_steps != steps: sigmas = sigmas[-(steps + 1):] + restart_segments = prepare_restart_segments(restart_info, real_model.model_sampling, sigmas) + sampler_wrapper = KSamplerRestartWrapper(sampler, real_model, restart_scheduler, restart_segments, seed, custom_noise, chunked=chunked_mode) latent = latent_image @@ -175,8 +197,20 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart_sigmas"], defaults=[None, 0, 0., 0., None])): + # sigmas: Sigmas for normal (outside of a restart segment) sampling. They start from after the previous PlanItem's steps + # if there is one or simply the beginning of sampling. + # k, s_min, s_max: This is the same as the restart segment definition. Set to 0 if there is no restart segment. + # restart_sigmas: Sigmas for the restart segment if it exists, otherwise None. + # Note: n_restart is not included as it can be calculated from the length of restart_sigmas. __slots__ = () + # Execute a plan item: runs sampling on the main sigmas, handles injecting noise for restarts + # as well as sampling the restart steps. + # sample: Function used sample sigmas. It takes x, a tensor with the sigmas to sample and + # the restart index (k) or -1 for sampling that isn't within a restart segment. + # get_noise_sampler: Return the noise sampler for restart segment noise injection. + # It takes x, and sigma_min, sigma_max (basically the same arguments as ComfyUI's + # BrownianTreeNoiseSampler class init function). @torch.no_grad() def execute(self, x, sample, get_noise_sampler): x = sample(x, self.sigmas, -1) @@ -190,6 +224,18 @@ class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart class KSamplerRestartWrapper: + # Some extra explanation for a couple of these arguments: + # + # chunked: + # When chunked is False, the sampling function is called step-by-step with only two sigmas at a time. + # When chunked is is True, the sampling function will be called with sigmas for multiple steps at a time. + # this means either the steps up to the next restart segment (or the end of sampling) or the steps within + # a restart segment. + # + # make_noise_sampler: + # If set to None, restart noise will just use torch.randn_like (gaussian) for noise generation. Otherwise + # this should contain a function that takes x, sigma_min, sigma_max, seed and returns a noise sampler + # function (which takes sigma, sigma_next) and returns a noisy tensor. 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 @@ -200,6 +246,9 @@ class KSamplerRestartWrapper: self.make_noise_sampler = make_noise_sampler self.chunked = chunked + # Builds a list of PlanItems and calculates the total number of steps. See the comments for PlanItem + # for more information about plans. + # Returns two values: the plan and the total steps. @torch.no_grad() def build_plan(self, sigmas, device): segments = round_restart_segments(sigmas, self.restart_segments) @@ -222,9 +271,15 @@ class KSamplerRestartWrapper: plan.append(PlanItem(sigmas[range_start:])) return plan, total_steps + # Dumps information about the plan to the console. It uses the normal plan execute + # logic. def explain_plan(self, plan, total_steps): + print(f"** Dumping restart sampling plan (total steps {total_steps}):") step = 0 last_kidx = -1 + # Instead of actually sampling, we just dump information about the steps. + # When kidx==-1 this is a normal step, otherwise kidx==0 is the first restart, + # kidx==1 is the second, etc. def do_sample(x, sigs, kidx=-1): nonlocal step, last_kidx rlabel = f"R{kidx+1:>3}" if kidx > last_kidx else " " @@ -240,11 +295,13 @@ class KSamplerRestartWrapper: step += chunk_size return x + # Stub function to satisfy PlanItem.execute def get_noise_sampler(*_args): return lambda *_args: 0.0 for pi in plan: pi.execute(0.0, do_sample, get_noise_sampler) + print("** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.") @torch.no_grad() @@ -254,11 +311,16 @@ class KSamplerRestartWrapper: seed = self.seed plan, self.total_steps = self.build_plan(sigmas, x.device) - self.explain_plan(plan, self.total_steps) + if VERBOSE: + self.explain_plan(plan, self.total_steps) def noise_sampler(*_args): return torch.randn_like(x) + # Passed to the PlanItem .execute method. Most of the time, self.make_noise_sampler + # is going to be None so this is just a wrapper for torch.randn_like. + # Otherwise we call the noise sampler factory and increment seed to ensure that restarts + # don't all use the same noise. def get_noise_sampler(x, s_min, s_max): nonlocal seed if not self.make_noise_sampler: @@ -276,6 +338,7 @@ class KSamplerRestartWrapper: if callback is not None: callback(x) + # Convenience function for code reuse. def sampler_function(x, sigs): return ksampler.sampler_function( model, x, sigs, *args, extra_args=extra_args, callback=callback_wrapper, disable=True, @@ -285,11 +348,15 @@ class KSamplerRestartWrapper: if isinstance(sigs, (list, tuple)): sigs = torch.tensor(sigs, device=x.device) if self.chunked or len(sigs) < 3: + # If running un chunked mode or there are already 2 or less sigmas, we can just + # pass the sigmas to the sampling function. return sampler_function(x, sigs) + # Otherwise we call the sampling function step by step on slices of 2 sigmas. for i in range(len(sigs)-1): x = sampler_function(x, sigs[i:i+2]) return x + # Execute the plan items in sequence. for pi in plan: x = pi.execute(x, do_sample, get_noise_sampler) From 17193ecdbf616b610b93c883e64ac2a8446e43ab Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 28 Mar 2024 02:00:34 -0600 Subject: [PATCH 08/10] Minor cleanups + add a few more comments Allow passing sigmas to main restart_sampling function --- nodes.py | 8 ++++---- restart_sampling.py | 15 ++++++++++----- 2 files changed, 14 insertions(+), 9 deletions(-) diff --git a/nodes.py b/nodes.py index b5914d4..70ed2f1 100644 --- a/nodes.py +++ b/nodes.py @@ -48,7 +48,7 @@ class KRestartSamplerSimple: FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, chunked_mode=False): + def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments): return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, scheduler, denoise=denoise) @@ -77,7 +77,7 @@ class KRestartSampler: FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler, chunked_mode=False): + def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler, chunked_mode=True): return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise, chunked_mode=chunked_mode) @@ -110,7 +110,7 @@ class KRestartSamplerAdv: FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler, chunked_mode=False): + def sample(self, model, add_noise, noise_seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, segments, restart_scheduler, chunked_mode=True): force_full_denoise = return_with_leftover_noise != "enable" disable_noise = add_noise == "disable" return restart_sampling(model, noise_seed, steps, cfg, sampler_name, 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, chunked_mode=chunked_mode) @@ -146,7 +146,7 @@ class KRestartSamplerCustom: FUNCTION = "sample" CATEGORY = "sampling" - 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, chunked_mode=False): + 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, chunked_mode=True): 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, chunked_mode=chunked_mode) diff --git a/restart_sampling.py b/restart_sampling.py index 143f718..64199f7 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -111,7 +111,7 @@ def calc_restart_steps(restart_segments): return restart_steps -def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, chunked_mode=True): +def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, negative, latent_image, restart_info, restart_scheduler, denoise=1.0, disable_noise=False, step_range=None, force_full_denoise=False, output_only=True, custom_noise=None, chunked_mode=True, sigmas=None): if isinstance(sampler, str): sampler = sampler_object(sampler) @@ -121,10 +121,13 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega real_model = real_model.model effective_steps = steps if step_range is not None or denoise > 0.9999 else int(steps / denoise) - sigmas = calc_sigmas(scheduler, effective_steps, - float(real_model.model_sampling.sigma_min), float(real_model.model_sampling.sigma_max), - real_model, model.load_device, - ) + if sigmas is None: + sigmas = calc_sigmas(scheduler, effective_steps, + float(real_model.model_sampling.sigma_min), float(real_model.model_sampling.sigma_max), + real_model, model.load_device, + ) + else: + sigmas = sigmas.detach().clone().to(model.load_device) if step_range is not None: start_step, last_step = step_range @@ -257,6 +260,7 @@ class KSamplerRestartWrapper: range_start = -1 for i in range(len(sigmas) - 1): if range_start == -1: + # Starting a new plan item - main sigmas start at the current index of i. range_start = i s_min = sigmas[i + 1].item() seg = segments.get(s_min) @@ -268,6 +272,7 @@ class KSamplerRestartWrapper: plan.append(PlanItem(sigmas[range_start:i+2], k, s_min, s_max, seg_sigmas[:-1])) range_start = -1 if range_start != -1: + # Include sigmas after the last restart segments in the plan. plan.append(PlanItem(sigmas[range_start:])) return plan, total_steps From f890b8bf622c88294078477229a11fe7425f8eaa Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 31 Mar 2024 16:04:15 -0600 Subject: [PATCH 09/10] Use correct value for a1111 mode t_max --- restart_sampling.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/restart_sampling.py b/restart_sampling.py index 64199f7..c8834f3 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -54,10 +54,10 @@ def prepare_restart_segments(restart_info, ms, sigmas): return [] if steps < 36: # Less than 36 steps - one restart with 9 steps. - restart_info = "[10,1,0.1,0.2]" + restart_info = "[10, 1, 0.1, 2.0]" else: # Otherwise two restarts with steps // 4 steps. - restart_info = f"[{(steps // 4) + 1}, 2, 0.1, 0.2]" + restart_info = f"[{(steps // 4) + 1}, 2, 0.1, 2.0]" try: restart_arrays = ast.literal_eval(f"[{restart_info}]") except SyntaxError as e: From 9c652149338280c5dd555ae03c504575cdba067b Mon Sep 17 00:00:00 2001 From: blepping Date: Mon, 1 Apr 2024 03:34:24 -0600 Subject: [PATCH 10/10] Attempt to get a1111 mode t_max calculation correct Improve sigmas output when dumping restart plan --- restart_sampling.py | 23 ++++++++++++++--------- 1 file changed, 14 insertions(+), 9 deletions(-) diff --git a/restart_sampling.py b/restart_sampling.py index c8834f3..876a02c 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -44,6 +44,7 @@ def prepare_restart_segments(restart_info, ms, sigmas): if restart_info == "": # No restarts. return [] + restart_arrays = None if restart_info == "default": restart_info = DEFAULT_SEGMENTS elif restart_info == "a1111": @@ -52,17 +53,19 @@ def prepare_restart_segments(restart_info, ms, sigmas): if steps < 20: # Less than 20 steps - no restarts. return [] + a1111_t_max = sigmas[int(torch.argmin(abs(sigmas - 2.0), dim=0))].item() if steps < 36: # Less than 36 steps - one restart with 9 steps. - restart_info = "[10, 1, 0.1, 2.0]" + restart_arrays = [[10, 1, 0.1, a1111_t_max]] else: # Otherwise two restarts with steps // 4 steps. - restart_info = f"[{(steps // 4) + 1}, 2, 0.1, 2.0]" - try: - restart_arrays = ast.literal_eval(f"[{restart_info}]") - except SyntaxError as e: - print("Ill-formed restart segments") - raise e + restart_arrays = [[(steps // 4) + 1, 2, 0.1, a1111_t_max]] + if restart_arrays is None: + try: + restart_arrays = ast.literal_eval(f"[{restart_info}]") + except SyntaxError as e: + print("Ill-formed restart segments") + raise e restart_segments = [] for arr in restart_arrays: if len(arr) != 4: @@ -279,6 +282,8 @@ class KSamplerRestartWrapper: # Dumps information about the plan to the console. It uses the normal plan execute # logic. def explain_plan(self, plan, total_steps): + def pretty_sigmas(sigmas): + return ", ".join(f"{sig:.4}" for sig in sigmas.tolist()) print(f"** Dumping restart sampling plan (total steps {total_steps}):") step = 0 last_kidx = -1 @@ -292,11 +297,11 @@ class KSamplerRestartWrapper: if not self.chunked: for i in range(len(sigs)-1): step += 1 - print(f"[{rlabel}] Step {step:>3}: {sigs[i:i+2]}") + print(f"[{rlabel}] Step {step:>3}: {pretty_sigmas(sigs[i:i+2])}") return x chunk_size = len(sigs) - 2 step += 1 - print(f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {sigs}") + print(f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigs)}") step += chunk_size return x