diff --git a/README.md b/README.md index eef07a2..98bd053 100644 --- a/README.md +++ b/README.md @@ -18,6 +18,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| @@ -39,6 +42,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 269c86f..70ed2f1 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}), } } @@ -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=True): + 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=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) + 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,11 +146,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, 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) - + 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, diff --git a/restart_sampling.py b/restart_sampling.py index 771730c..876a02c 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -1,4 +1,6 @@ import ast +from collections import namedtuple +import os import warnings import torch from tqdm.auto import trange @@ -9,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: @@ -33,12 +39,33 @@ def resolve_t_value(val, ms): raise ValueError("bad t_min or t_max value") -def prepare_restart_segments(restart_info, ms): - try: - restart_arrays = ast.literal_eval(f"[{restart_info}]") - except SyntaxError as e: - print("Ill-formed restart segments") - raise e +def prepare_restart_segments(restart_info, ms, sigmas): + restart_info = restart_info.strip().lower() + if restart_info == "": + # No restarts. + return [] + restart_arrays = None + 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 [] + 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_arrays = [[10, 1, 0.1, a1111_t_max]] + else: + # Otherwise two restarts with steps // 4 steps. + 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: @@ -87,23 +114,23 @@ 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, chunked_mode=True, sigmas=None): 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), - 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 @@ -117,8 +144,9 @@ 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) + 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 latent_image = latent["samples"] @@ -148,7 +176,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 @@ -174,49 +202,172 @@ 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])): + # 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) + 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, total_steps): + # 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 self.restart_scheduler = restart_scheduler self.restart_segments = restart_segments - self.total_steps = total_steps + self.total_steps = 0 + self.seed = seed + 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) + total_steps = len(sigmas) - 1 + calc_restart_steps(segments) + plan = [] + 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) + 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=device) + 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 + + # 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 + # 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 " " + last_kidx = kidx + if not self.chunked: + for i in range(len(sigs)-1): + step += 1 + 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}: {pretty_sigmas(sigs)}") + 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() def ksampler_restart_wrapper(self, model, x, sigmas, *args, extra_args=None, callback=None, disable=None, **kwargs): ksampler = self.ksampler - segments = round_restart_segments(sigmas, self.restart_segments) - self.total_steps[0] = len(sigmas) - 1 + calc_restart_steps(segments) step = 0 + seed = self.seed + plan, self.total_steps = self.build_plan(sigmas, x.device) - 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) + 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: + 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): + nonlocal step step += 1 - 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) - for _ in range(k): - x += torch.randn_like(x) * (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 + pbar.update(1) + x["i"] = step + 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, + **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: + # 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) + return x