diff --git a/nodes.py b/nodes.py index 7dddf11..a46eee2 100644 --- a/nodes.py +++ b/nodes.py @@ -1,13 +1,20 @@ +import os + import comfy -import torch from .restart_sampling import ( DEFAULT_SEGMENTS, SCHEDULER_MAPPING, + VERBOSE, RestartPlan, + RestartSampler, restart_sampling, ) +INCLUDE_SELFTEST = ( + os.environ.get("COMFYUI_RESTART_SAMPLING_SELFTEST", "").strip() == "1" +) + def get_supported_samplers(): samplers = comfy.samplers.KSampler.SAMPLERS.copy() @@ -302,7 +309,7 @@ class KRestartSamplerCustom: ) -class RestartScheduler: +class RestartSchedulerNode: @classmethod def INPUT_TYPES(cls): return { @@ -336,8 +343,6 @@ class RestartScheduler: denoise, sigmas_opt=None, ): - # RestartPlan.self_test(model, max_steps=200) - plan = RestartPlan( model, steps, @@ -347,11 +352,12 @@ class RestartScheduler: denoise, sigmas=sigmas_opt, ) - plan.explain(chunked=True) + if VERBOSE: + plan.explain(chunked=True) return (plan.sigmas(),) -class RestartSampler: +class RestartSamplerNode: @classmethod def INPUT_TYPES(cls): return { @@ -366,39 +372,16 @@ class RestartSampler: CATEGORY = "sampling/custom_sampling/samplers" def go(self, sampler, chunked_mode): - wrapped = comfy.samplers.KSAMPLER( - lambda *args, **kwargs: self.sampler_function( - sampler, - chunked_mode, - *args, - **kwargs, - ), - extra_options=sampler.extra_options, + restart_options = { + "restart_chunked": chunked_mode, + "restart_wrapped_sampler": sampler, + } + restart_sampler = comfy.samplers.KSAMPLER( + RestartSampler.sampler_function, + extra_options=sampler.extra_options | restart_options, inpaint_options=sampler.inpaint_options, ) - return (wrapped,) - - @staticmethod - @torch.no_grad() - def sampler_function( - wrapped, - chunked, - model, - x, - sigmas, - *args: list, - **kwargs: dict, - ) -> torch.Tensor: - plan = RestartPlan.from_sigmas(sigmas) - return plan.sample( - wrapped, - model, - x, - sigmas, - *args, - restart_chunked=chunked, - **kwargs, - ) + return (restart_sampler,) NODE_CLASS_MAPPINGS = { @@ -406,8 +389,8 @@ NODE_CLASS_MAPPINGS = { "KRestartSampler": KRestartSampler, "KRestartSamplerAdv": KRestartSamplerAdv, "KRestartSamplerCustom": KRestartSamplerCustom, - "RestartScheduler": RestartScheduler, - "RestartSampler": RestartSampler, + "RestartScheduler": RestartSchedulerNode, + "RestartSampler": RestartSamplerNode, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -416,3 +399,28 @@ NODE_DISPLAY_NAME_MAPPINGS = { "KRestartSamplerAdv": "KSampler With Restarts (Advanced)", "KRestartSamplerCustom": "KSampler With Restarts (Custom)", } + +if os.environ.get("COMFYUI_RESTART_SAMPLING_SELFTEST", "").strip() == "1": + + class RestartSelfTestNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "enabled": ("BOOLEAN", {"default": True}), + "min_steps": ("INT", {"default": 2, "min": 0}), + "max_steps": ("INT", {"default": 100, "min": 2}), + }, + } + + RETURN_TYPES = ("BOOLEAN",) + FUNCTION = "go" + CATEGORY = "sampling/custom_sampling/samplers" + + def go(self, model, enabled=True, min_steps=2, max_steps=100): + if enabled: + RestartPlan.self_test(model, min_steps=min_steps, max_steps=max_steps) + return (True,) + + NODE_CLASS_MAPPINGS["RestartSelfTest"] = RestartSelfTestNode diff --git a/restart_sampling.py b/restart_sampling.py index 53bd463..c9e1cc9 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -190,19 +190,16 @@ def restart_sampling( force_full_denoise=force_full_denoise, sigmas=sigmas, ) - plan = plan.to(model.load_device) - ### UNCOMMENT TO RUN SELF TEST - # plan.self_test( - # model, - # min_steps=2, - # schedules=SCHEDULER_MAPPING.keys(), - # # schedules=("simple_test",), - # restart_schedules=SCHEDULER_MAPPING.keys(), - # ) - sigmas = plan.sigmas() + + if VERBOSE: + plan.explain(chunked_mode) + + total_steps = plan.total_steps + sigmas = plan.sigmas().to(model.load_device) latent = latent_image latent_image = latent["samples"] + if disable_noise: torch.manual_seed( seed, @@ -222,32 +219,27 @@ def restart_sampling( noise_mask = latent["noise_mask"] x0_output = {} - callback = latent_preview.prepare_callback( - model, - sigmas.shape[-1] - 1, - x0_output, - ) + callback = latent_preview.prepare_callback(model, plan.total_steps, x0_output) disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED + restart_options = { + "restart_chunked": chunked_mode, + "restart_wrapped_sampler": sampler, + "restart_custom_noise": custom_noise, + } + ksampler = KSAMPLER( - lambda *args, **kwargs: plan.sample( - sampler, - *args, - restart_chunked=chunked_mode, - restart_make_noise_sampler=custom_noise, - restart_seed=seed, - **kwargs, - ), - extra_options=sampler.extra_options | {}, + RestartSampler.sampler_function, + extra_options=sampler.extra_options | restart_options, inpaint_options=sampler.inpaint_options | {}, ) # Add the additional steps to the progress bar pbar_update_absolute = ProgressBar.update_absolute - def pbar_update_absolute_wrapper(self, value, total=None, preview=None): - pbar_update_absolute(self, value, plan.total_steps, preview) + def pbar_update_absolute_wrapper(self, value, total=None, preview=None): # noqa: ARG001 + pbar_update_absolute(self, value, total_steps, preview) ProgressBar.update_absolute = pbar_update_absolute_wrapper @@ -352,35 +344,6 @@ class PlanItem( def s_max(self): return None if self.k < 1 else self.restart_sigmas[0].item() - # 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, skip_normal=False, next_pi=None): - if not skip_normal: - 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 - ) - if next_pi and kidx == self.k - 1: - sigmas = torch.cat((self.restart_sigmas[:-1], next_pi.sigmas)).to( - self.restart_sigmas.device, - ) - print("COMBINE", sigmas) - else: - sigmas = self.restart_sigmas - x = sample(x, sigmas, kidx) - return x - class RestartPlan: def __init__( @@ -395,10 +358,6 @@ class RestartPlan: force_full_denoise=False, sigmas=None, ): - # comfy.model_management.load_models_gpu([model]) - # real_model = model - # while hasattr(real_model, "model"): - # real_model = real_model.model ms = model.get_model_object("model_sampling") effective_steps = ( @@ -430,8 +389,6 @@ class RestartPlan: elif effective_steps != steps: sigmas = sigmas[-(steps + 1) :] - self.plain_sigmas = sigmas - restart_segments = prepare_restart_segments(restart_info, ms, sigmas) self.plan, self.total_steps = self.build_plan_items( model.model, @@ -444,9 +401,6 @@ class RestartPlan: def __repr__(self) -> str: return f"" - def __len__(self) -> int: - return self.total_steps - # 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. @@ -459,6 +413,7 @@ class RestartPlan: sigmas, device, ) -> tuple[list, int]: + model_sigma_min = float(model.model_sampling.sigma_min) segments = round_restart_segments(sigmas, restart_segments) plan = [] range_start = -1 @@ -473,14 +428,14 @@ class RestartPlan: s_max, k, n_restart = seg["t_max"], seg["k"], seg["n"] if k < 1 or n_restart < 2: continue + if s_max <= model_sigma_min: + errstr = f"Restart: Invalid restart segment t_max {s_max:.05} <= model minimum sigma {model_sigma_min:.05}" + raise ValueError(errstr) normal_sigmas = sigmas[range_start : i + 2] - effsmin = max(float(model.model_sampling.sigma_min), sigmas[i + 1]) - if effsmin >= s_max: - continue restart_sigmas = calc_sigmas( restart_scheduler, n_restart, - effsmin, + max(model_sigma_min, sigmas[i + 1]), s_max, model, device=device, @@ -496,199 +451,176 @@ class RestartPlan: plan.append(PlanItem(sigmas[range_start:])) return plan, sum(pi.total_steps for pi in plan) - @classmethod - def from_sigmas(cls, sigmas, threshold=1e-06): - def get_normal_segment(sigmas): - # A normal segment ends when we either reach the end of the list or - # encounter a sigma higher than the previous. - last_sigma = sigmas[0] - for idx in range(1, len(sigmas)): - sigma = sigmas[idx] - if last_sigma - sigma < threshold: - return sigmas[:idx] - last_sigma = sigma - return sigmas - - def get_restart_segment(sigmas, s_min): - # s_min here is the last sigma of the previous normal segment. A restart segment - # ends when: - # 1. We reach the end of the list, or - # 2. We hit a sigma greater or equal to the last sigma, or - # 3. We hit a sigma less than s_min - last_sigma = sigmas[0] - for idx in range(1, len(sigmas)): - sigma = sigmas[idx] - if last_sigma - sigma < -threshold: - return sigmas[:idx] - if sigma <= s_min: - return sigmas[: idx + 1] - last_sigma = sigma - raise ValueError("Unexpected end of sigmas in a restart segment") - - plain_sigmas = sigmas.clone().detach().cpu() - plan = [] - while len(sigmas) > 0: - # Get the normal segment - a restart segment can never be first. - normal_sigmas = get_normal_segment(sigmas) - nslen = len(normal_sigmas) - if nslen < 2: - print(sigmas) - raise ValueError( - "Encountered invalid normal segment rebuilding sigmas: too short", - ) - sigmas = sigmas[nslen:] - if len(sigmas) == 0: - # No restart segments follow the normal segment so we're done. - plan.append(PlanItem(normal_sigmas)) - break - # If we're here there has to be a restart segment; get it. - restart_sigmas = get_restart_segment(sigmas, normal_sigmas[-1]) - rslen = len(restart_sigmas) - if rslen < 2: - raise ValueError( - "Encountered invalid normal segment rebuilding sigmas: too short", - ) - sigmas = sigmas[rslen:] - k = 1 - # The restart segment may be repeated multiple times. If so, count the - # repeats and trim the sigmas list. - while len(sigmas) > 0 and torch.equal(sigmas[:rslen], restart_sigmas): - k += 1 - sigmas = sigmas[rslen:] - plan.append(PlanItem(normal_sigmas, k, restart_sigmas)) - obj = cls.__new__(cls) - obj.plan = plan - obj.total_steps = sum(pi.total_steps for pi in plan) - obj.plain_sigmas = plain_sigmas - return obj - def sigmas(self) -> torch.Tensor: def sigmas_generator(): - for pi in self.plan: - yield pi.sigmas.cpu() - for _ in range(pi.k): - yield pi.restart_sigmas.cpu() + skip = False + plan = self.plan + planlen = len(plan) + for idx in range(planlen): + pi = plan[idx] + nextpi = None if idx == planlen - 1 else plan[idx + 1] + if not skip: + yield pi.sigmas + skip = False + if pi.k == 0: + continue + for _ in range(pi.k - 1): + yield pi.restart_sigmas + if nextpi is None: + yield pi.restart_sigmas + continue + skip = True + yield pi.restart_sigmas[:-1] + yield nextpi.sigmas return torch.flatten(torch.cat(tuple(sigmas_generator()))) - def to(self, device): - obj = self.__class__.__new__(self.__class__) - obj.plain_sigmas = self.plain_sigmas.to(device) - obj.total_steps = self.total_steps - items = obj.plan = [] - for pi in self.plan: - sigmas = pi.sigmas.to(device) - if pi.k < 1: - items.append(PlanItem(sigmas)) - continue - items.append(PlanItem(sigmas, pi.k, pi.restart_sigmas.to(device))) - return obj - # Dumps information about the plan to the console. It uses the normal plan execute # logic. def explain(self, chunked=True): def pretty_sigmas(sigmas): return ", ".join(f"{sig:.4}" for sig in sigmas.tolist()) - print(f"** Dumping restart sampling plan (total steps {self.total_steps}):") - for pi in self.plan: - print( - f"\n{pi.sigmas[-1].item():.04} .. {pi.sigmas[0].item():.04} ({len(pi.sigmas)})", - ) - if pi.k > 0: + def dump_steps(step, sigmas, restart=0): + rlabel = f"R{restart:>3}" if restart > 0 else " " + if chunked: + chunk_size = len(sigmas) - 2 + step += 1 print( - f" {pi.restart_sigmas[-1].item():.04} ({pi.s_min:.04}) .. {pi.restart_sigmas[0].item():.04} ({pi.s_max:.04}): k={pi.k} ({len(pi.restart_sigmas)})", + f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigmas)}", ) + step += chunk_size + return step + for i in range(len(sigmas) - 1): + step += 1 + print(f"[{rlabel}] Step {step:>3}: {pretty_sigmas(sigmas[i:i+2])}") + return step + + print(f"** Dumping restart sampling plan (total steps {self.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 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: list): - return lambda *_args: 0.0 - for pi in self.plan: - pi.execute(0.0, do_sample, get_noise_sampler) + step = dump_steps(step, pi.sigmas) + for kidx in range(pi.k): + step = dump_steps(step, pi.restart_sigmas, kidx + 1) print( "** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.", ) + @staticmethod + def self_test( + model, + schedules=None, + restart_schedules=None, + segments=None, + min_steps=2, + max_steps=100, + ) -> None: + if schedules is None: + schedules = SCHEDULER_MAPPING.keys() + if restart_schedules is None: + restart_schedules = SCHEDULER_MAPPING.keys() + if segments is None: + segments = ("default", "a1111") + for schname in schedules: + for rschname in restart_schedules: + for tsegs in segments: + print( + f"--- Test: {min_steps}..{max_steps} steps, schedules {schname}/{rschname}, segments {tsegs}", + ) + for tsteps in range(min_steps, max_steps + 1): + label = f"** {tsteps:03}: {schname}, {rschname}, {tsegs}:" + try: + _plan = RestartPlan( + model, + tsteps, + schname, + tsegs, + rschname, + 1.0, + ) + except ValueError as err: + print(f"{label}\n\t!! FAIL: {err}") + raise + continue + print("\n|| Done test") + + +class RestartSampler: + @staticmethod + def get_segment(sigmas: torch.Tensor) -> torch.Tensor: + # A normal segment ends when we either reach the end of the list or + # encounter a sigma higher than the previous. + last_sigma = sigmas[0] + for idx in range(1, len(sigmas)): + sigma = sigmas[idx] + if sigma > last_sigma: + return sigmas[:idx] + last_sigma = sigma + return sigmas + + @classmethod + def split_sigmas(cls, sigmas): + prev_seg = None + while len(sigmas) > 1: + seg = cls.get_segment(sigmas) + sigmas = sigmas[len(seg) :] + if prev_seg is not None and seg[0] > prev_seg[-1]: + print( + f"CALC NOISE: min={prev_seg[-1].item():.04}, max={seg[0].item():.04}", + ) + noise_scale = ((seg[0] ** 2 - prev_seg[-1] ** 2) ** 0.5).item() + else: + noise_scale = 0.0 + prev_seg = seg + yield (noise_scale, seg) + # 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. + # restart_chunked: + # When False, the sampling function is called step-by-step with only two sigmas at a time. + # When 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: + # restart_custom_noise: # 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. + @classmethod @torch.no_grad() - def sample( - self, - ksampler, + def sampler_function( + cls, model, x, - _sigmas, + sigmas, *args: list, + restart_wrapped_sampler=None, restart_chunked=True, - restart_make_noise_sampler=None, - restart_seed=None, - extra_args=None, + restart_custom_noise=None, callback=None, disable=None, **kwargs: dict, - ): + ) -> torch.Tensor: + if not restart_wrapped_sampler: + raise ValueError("RestartSampler: missing restart_sampler option!") + + def restart_noise(x, _s_min, _s_max, _seed): + return lambda _s, _sn: torch.randn_like(x) + + seed = (kwargs.get("extra_args", {}) or {}).get("seed", 42) + if restart_custom_noise is not None: + restart_noise = restart_custom_noise + + sampler = restart_wrapped_sampler.sampler_function + + print("SAMPLING", sigmas) + total_steps = len(sigmas - 1) step = 0 - if restart_seed is None: - seed = (extra_args or {}).get("seed", 42) - else: - seed = restart_seed - plan = self.plan - - if VERBOSE: - self.explain(restart_chunked) - - def noise_sampler(*_args: list): - 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 restart_make_noise_sampler: - return noise_sampler - result = restart_make_noise_sampler(x, s_min, s_max, seed) - seed += 1 - return result - - with trange(self.total_steps, disable=disable) as pbar: + noise_count = 0 + with trange(total_steps, disable=disable) as pbar: last_cb_sigma = None - def callback_wrapper(cb_state): + def cb_wrapper(cb_state): nonlocal step, last_cb_sigma curr_sigma = cb_state.get("sigma") curr_sigma = ( @@ -706,117 +638,38 @@ class RestartPlan: if callback is not None: callback(cb_state) - # 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 restart_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. - skip = False - for idx in range(len(plan)): - pi = plan[idx] - nextpi = plan[idx + 1] if idx < len(plan) - 1 else None - x = pi.execute( - x, - do_sample, - get_noise_sampler, - skip_normal=skip, - next_pi=nextpi, - ) - skip = pi.k > 0 and len(pi.restart_sigmas) > 2 and nextpi is not None - return x - - @staticmethod - def self_test( - model, - schedules=None, - restart_schedules=None, - segments=None, - min_steps=2, - max_steps=100, - ) -> None: - if schedules is None: - schedules = SCHEDULER_MAPPING.keys() - {"simple_test"} - if restart_schedules is None: - restart_schedules = SCHEDULER_MAPPING.keys() - {"simple_test"} - if segments is None: - segments = ("default", "a1111") - for schname in schedules: - for rschname in restart_schedules: - for tsegs in segments: - print( - f"--- Test: {min_steps}..{max_steps} steps, schedules {schname}/{rschname}, segments {tsegs}", + for noise_scale, chunk_sigmas in cls.split_sigmas(sigmas): + print(f"CHUNK: noise={noise_scale:.04}, sigmas={chunk_sigmas}") + if noise_scale != 0: + x += ( + restart_noise( + x, + chunk_sigmas[-1], + chunk_sigmas[0], + seed + noise_count, + )(chunk_sigmas[0], chunk_sigmas[-1]) + * noise_scale ) - for tsteps in range(min_steps, max_steps + 1): - label = f"** {tsteps:03}: {schname}, {rschname}, {tsegs}:" - try: - p1 = RestartPlan( - model, - tsteps, - schname, - tsegs, - rschname, - 1.0, - ) - except ValueError as err: - print(f"{label}\n\t!! FAIL: {err}") - raise - continue - try: - p2 = RestartPlan.from_sigmas(p1.sigmas()) - except ValueError: - print(label) - p1.explain(chunked=True) - raise - fail = None - if len(p1) != len(p2): - fail = "steps" - if not fail: - for idx in range(len(p1.plan)): - pi1, pi2 = p1.plan[idx], p2.plan[idx] - if not torch.equal( - torch.round(pi1.sigmas, decimals=5), - torch.round(pi2.sigmas, decimals=5), - ): - fail = "normal" - break - if pi1.k != pi2.k: - fail = "k" - break - if pi1.k < 1: - continue - if not torch.equal( - torch.round(pi1.restart_sigmas, decimals=5), - torch.round(pi2.restart_sigmas, decimals=5), - ): - fail = "restart" - break - - if fail: - print(label) - print("!!!", fail) - p1.explain() - print("====") - p2.explain() - raise ValueError("Failed rebuilding restart plan") - print("\n|| Done test") + noise_count += 1 + if restart_chunked: + x = sampler( + model, + x, + chunk_sigmas, + *args, + callback=cb_wrapper, + disable=True, + **kwargs, + ) + continue + for i in range(len(chunk_sigmas) - 1): + x = sampler( + model, + x, + chunk_sigmas[i : i + 2], + *args, + callback=cb_wrapper, + disable=True, + **kwargs, + ) + return x