From e2e25e7a2870e0cacabc01e28106df2d06acea42 Mon Sep 17 00:00:00 2001 From: blepping Date: Wed, 27 Mar 2024 05:45:06 -0600 Subject: [PATCH] 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)