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)