diff --git a/nodes.py b/nodes.py index 6e151af..7dddf11 100644 --- a/nodes.py +++ b/nodes.py @@ -1,12 +1,10 @@ import comfy import torch -from . import restart_sampling as restart from .restart_sampling import ( DEFAULT_SEGMENTS, SCHEDULER_MAPPING, - KSamplerRestartWrapper, - rebuild_plan, + RestartPlan, restart_sampling, ) @@ -328,13 +326,6 @@ class RestartScheduler: FUNCTION = "go" CATEGORY = "sampling/custom_sampling/schedulers" - @staticmethod - def plan_sigmas(plan): # noqa: ANN205 - for pi in plan: - yield pi.sigmas - for _ in range(pi.k): - yield pi.restart_sigmas - def go( self, model, @@ -345,38 +336,19 @@ class RestartScheduler: denoise, sigmas_opt=None, ): - ms = model.get_model_object("model_sampling") - if sigmas_opt is None or len(sigmas_opt) < 2: - total_steps = steps - if denoise < 1.0: - if denoise <= 0.0: - return (torch.FloatTensor([]),) - total_steps = int(steps / denoise) + # RestartPlan.self_test(model, max_steps=200) - sigmas = restart.calc_sigmas( - scheduler, - total_steps, - float(ms.sigma_min), - float(ms.sigma_max), - model.model, - "cpu", - ) - sigmas = sigmas[-(steps + 1) :] - else: - sigmas = sigmas_opt - prepared_segments = restart.prepare_restart_segments(segments, ms, sigmas) - plan, restart_steps = restart.build_plan( - model.model, - prepared_segments, + plan = RestartPlan( + model, + steps, + scheduler, + segments, restart_scheduler, - sigmas, - "cpu", + denoise, + sigmas=sigmas_opt, ) - if restart.VERBOSE: - restart.explain_plan(plan, restart_steps, chunked=True) - restart_sigmas = torch.flatten(torch.cat(tuple(self.plan_sigmas(plan)))) - print("MADE SIGMAS", restart_sigmas) - return (restart_sigmas,) + plan.explain(chunked=True) + return (plan.sigmas(),) class RestartSampler: @@ -408,19 +380,25 @@ class RestartSampler: @staticmethod @torch.no_grad() - def sampler_function(wrapped, chunked, model, x, sigmas, *args, **kwargs): - plan, total_steps = rebuild_plan(sigmas) - print("Rebuilt", total_steps, plan) - seed = kwargs.get("extra_args", {}).get("seed") - rw = KSamplerRestartWrapper( + def sampler_function( + wrapped, + chunked, + model, + x, + sigmas, + *args: list, + **kwargs: dict, + ) -> torch.Tensor: + plan = RestartPlan.from_sigmas(sigmas) + return plan.sample( wrapped, - None, - None, - None, - seed, - chunked=chunked, + model, + x, + sigmas, + *args, + restart_chunked=chunked, + **kwargs, ) - return rw.sample_plan(plan, total_steps, model, x, sigmas, *args, **kwargs) NODE_CLASS_MAPPINGS = { diff --git a/restart_sampling.py b/restart_sampling.py index bfa7c75..4bfa0e8 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import ast import os import warnings @@ -42,14 +44,7 @@ def resolve_t_value(val, ms): 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": + def get_a1111_segment(): # Emulate A1111 WebUI's restart sampler behavior. steps = len(sigmas) - 1 if steps < 20: @@ -58,20 +53,46 @@ def prepare_restart_segments(restart_info, ms, sigmas): 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]] + return [10, 1, 0.1, a1111_t_max] + # Otherwise two restarts with steps // 4 steps. + return [(steps // 4) + 1, 2, 0.1, a1111_t_max] + + 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": + restart_arrays = [get_a1111_segment()] + if restart_arrays == [[]]: + return [] if restart_arrays is None: try: restart_arrays = ast.literal_eval(f"[{restart_info}]") except SyntaxError: print("Ill-formed restart segments") raise + temp = [] + default_segments = ast.literal_eval(DEFAULT_SEGMENTS) + for idx in range(len(restart_arrays)): + item = restart_arrays[idx] + if not isinstance(item, str): + temp.append(item) + continue + preset = item.strip().lower() + if preset == "default": + temp += default_segments + elif preset == "a1111": + temp.append(get_a1111_segment()) + else: + raise ValueError("Ill-formed restart segment") + restart_arrays = temp restart_segments = [] for arr in restart_arrays: - if len(arr) != 4: - raise ValueError("Restart segment must have 4 values") + if not isinstance(arr, (list, tuple)) or len(arr) != 4: + raise ValueError("Restart segment must be a list with 4 values") n_restart, k, val_min, val_max = arr n_restart, k = int(n_restart), int(k) t_min = resolve_t_value(val_min, ms) @@ -158,53 +179,19 @@ def restart_sampling( 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 - - effective_steps = ( - steps if step_range is not None or denoise > 0.9999 else int(steps / denoise) - ) - 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 - - if last_step < (len(sigmas) - 1): - sigmas = sigmas[: last_step + 1] - if force_full_denoise: - sigmas[-1] = 0 - - if start_step < (len(sigmas) - 1): - sigmas = sigmas[start_step:] - elif effective_steps != steps: - sigmas = sigmas[-(steps + 1) :] - - restart_segments = prepare_restart_segments( + plan = RestartPlan( + model, + steps, + scheduler, restart_info, - real_model.model_sampling, - sigmas, - ) - - sampler_wrapper = KSamplerRestartWrapper( - sampler, - real_model, restart_scheduler, - restart_segments, - seed, - custom_noise, - chunked=chunked_mode, + denoise=denoise, + step_range=step_range, + force_full_denoise=force_full_denoise, + sigmas=sigmas, ) + plan = plan.to(model.load_device) + sigmas = plan.sigmas() latent = latent_image latent_image = latent["samples"] @@ -227,12 +214,23 @@ 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, + sigmas.shape[-1] - 1, + x0_output, + ) disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED - sampler = KSAMPLER( - sampler_wrapper.ksampler_restart_wrapper, + 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 | {}, inpaint_options=sampler.inpaint_options | {}, ) @@ -241,7 +239,7 @@ def restart_sampling( pbar_update_absolute = ProgressBar.update_absolute def pbar_update_absolute_wrapper(self, value, total=None, preview=None): - pbar_update_absolute(self, value, sampler_wrapper.total_steps, preview) + pbar_update_absolute(self, value, plan.total_steps, preview) ProgressBar.update_absolute = pbar_update_absolute_wrapper @@ -250,7 +248,7 @@ def restart_sampling( model, noise, cfg, - sampler, + ksampler, sigmas, positive, negative, @@ -285,6 +283,42 @@ class PlanItem( defaults=[None, 0, 0.0, 0.0, None], ), ): + def __new__(cls, *args: list, **kwargs: dict): + threshold = 1e-06 + obj = super().__new__(cls, *args, **kwargs) + if len(obj.sigmas) < 2: + raise ValueError("PlanItem: invalid normal sigmas: too short") + if obj.k < 1: + return obj + if len(obj.restart_sigmas) < 2: + raise ValueError("PlanItem: invalid restart sigmas: too short") + if obj.s_min >= obj.s_max: + raise ValueError("PlanItem: invalid min/max: min >= max") + # if obj.sigmas[-1] >= obj.restart_sigmas[0]: + if obj.sigmas[-1] - obj.restart_sigmas[0] > threshold: + raise ValueError( + "PlanItem: invalid sigmas: last normal sigma >= first restart sigma", + ) + # if obj.restart_sigmas[-1] < obj.sigmas[-1]: + if obj.sigmas[-1] - obj.restart_sigmas[-1] > 1e-02: # threshold: + errstr = ( + f"PlanItem: invalid sigmas: last restart sigma {obj.restart_sigmas[-1]} < last normal sigma {obj.sigmas[-1]}", + ) + raise ValueError(errstr) + t = obj.sigmas.sort(descending=True, stable=True)[0].unique_consecutive() + if not torch.equal(obj.sigmas, t): + raise ValueError( + "PlanItem: invalid normal sigmas: out of order or contains duplicates", + ) + t = obj.restart_sigmas.sort(descending=True, stable=True)[ + 0 + ].unique_consecutive() + if not torch.equal(obj.restart_sigmas, t): + raise ValueError( + "PlanItem: invalid restart sigmas: out of order or contains duplicates", + ) + return obj + # 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. @@ -314,141 +348,281 @@ class PlanItem( return x -# 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(model, restart_segments, restart_scheduler, sigmas, device): - segments = round_restart_segments(sigmas, 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( - restart_scheduler, - n_restart, - s_min, - s_max, - model, - device=device, +class RestartPlan: + def __init__( + self, + model, + steps, + scheduler, + restart_info, + restart_scheduler, + denoise=1.0, + step_range=None, + 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 + + effective_steps = ( + steps + if step_range is not None or denoise > 0.9999 + else int(steps / denoise) ) - 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 - - -def rebuild_plan(sigmas): - def get_normal_segment(sigmas): - last_sigma = None - for idx in range(len(sigmas)): - sigma = sigmas[idx] - if last_sigma is not None and sigma >= last_sigma: - return sigmas[:idx] - last_sigma = sigma - return sigmas - - def get_restart_segment(sigmas, s_min): - last_sigma = None - for idx in range(len(sigmas)): - sigma = sigmas[idx] - if (last_sigma is not None and sigma >= last_sigma) or sigma < s_min: - return sigmas[:idx] - last_sigma = sigma - raise ValueError("Unexpected end of sigmas in a restart segment") - - plan = [] - total_steps = 0 - while len(sigmas) > 0: - normal_sigmas = get_normal_segment(sigmas) - nslen = len(normal_sigmas) - sigmas = sigmas[nslen:] - total_steps += nslen - 1 - if len(sigmas) == 0: - plan.append(PlanItem(normal_sigmas)) - break - restart_sigmas = get_restart_segment(sigmas, normal_sigmas[-1]) - rslen = len(restart_sigmas) - sigmas = sigmas[rslen:] - k = 1 - while len(sigmas) > 0 and torch.equal(sigmas[:rslen], restart_sigmas): - k += 1 - sigmas = sigmas[rslen:] - total_steps += (rslen - 1) * k - plan.append( - PlanItem( - normal_sigmas, - k, - normal_sigmas[-1], - restart_sigmas[0], - restart_sigmas, - ), - ) - return plan, total_steps - - -# Dumps information about the plan to the console. It uses the normal plan execute -# logic. -def explain_plan(plan, total_steps, chunked=True): - def pretty_sigmas(sigmas): - return ", ".join(f"{sig:.4}" for sig in sigmas.tolist()) - - print(f"** Dumping restart sampling plan (total steps {total_steps}):") - for pi in plan: - print( - f"\n{pi.sigmas[-1].item():.04} .. {pi.sigmas[0].item():.04} ({len(pi.sigmas)})", - ) - if pi.k > 0: - 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)})", + 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, + "cpu", + # model.load_device, ) - step = 0 - last_kidx = -1 + else: + sigmas = sigmas.detach().cpu().clone() + if step_range is not None: + start_step, last_step = step_range - # 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)}", + if last_step < (len(sigmas) - 1): + sigmas = sigmas[: last_step + 1] + if force_full_denoise: + sigmas[-1] = 0 + + if start_step < (len(sigmas) - 1): + sigmas = sigmas[start_step:] + elif effective_steps != steps: + sigmas = sigmas[-(steps + 1) :] + + self.plain_sigmas = sigmas + + restart_segments = prepare_restart_segments( + restart_info, + real_model.model_sampling, + sigmas, + ) + self.plan, self.total_steps = self.build_plan_items( + model.model, + restart_segments, + restart_scheduler, + sigmas, + "cpu", ) - step += chunk_size - return x - # Stub function to satisfy PlanItem.execute - def get_noise_sampler(*_args): - return lambda *_args: 0.0 + def __repr__(self) -> str: + return f"" - 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.", - ) + 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. + @staticmethod + @torch.no_grad() + def build_plan_items( + model, + restart_segments, + restart_scheduler, + sigmas, + device, + ) -> tuple[list, int]: + segments = round_restart_segments(sigmas, 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( + restart_scheduler, + n_restart, + s_min, + s_max, + 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 + + @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] + if s_min - last_sigma > threshold: + return sigmas[:2] + for idx in range(1, len(sigmas)): + sigma = sigmas[idx] + # sigma > last_sigma + if last_sigma - sigma < threshold: + return sigmas[:idx] + + # sigma < s_min + if s_min - sigma > -threshold: + # TODO: Document this part + if idx < len(sigmas) - 2 and sigmas[idx + 1] - s_min < -threshold: + return sigmas[:idx] + return sigmas[: idx + 1] + last_sigma = sigma + raise ValueError("Unexpected end of sigmas in a restart segment") + + plain_sigmas = sigmas.detach().cpu().clone() + plan = [] + total_steps = 0 + 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:] + total_steps += nslen - 1 + 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: + print(restart_sigmas) + 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:] + total_steps += (rslen - 1) * k + plan.append( + PlanItem( + normal_sigmas, + k, + normal_sigmas[-1], + restart_sigmas[0], + restart_sigmas, + ), + ) + obj = cls.__new__(cls) + obj.plan = plan + obj.total_steps = total_steps + 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() + + 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.s_min, + pi.s_max, + 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: + 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)})", + ) + 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) + print( + "** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.", + ) -class KSamplerRestartWrapper: # Some extra explanation for a couple of these arguments: # # chunked: @@ -461,48 +635,33 @@ class KSamplerRestartWrapper: # 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 = 0 - self.seed = seed - self.make_noise_sampler = make_noise_sampler - self.chunked = chunked - @torch.no_grad() - def sample_plan( + def sample( self, - plan, - total_steps, + ksampler, model, x, - sigmas, - *args, + _sigmas, + *args: list, + restart_chunked=True, + restart_make_noise_sampler=None, + restart_seed=None, extra_args=None, callback=None, disable=None, - **kwargs, + **kwargs: dict, ): - self.total_steps = total_steps - ksampler = self.ksampler step = 0 - seed = self.seed + if restart_seed is None: + seed = (extra_args or {}).get("seed", 42) + else: + seed = restart_seed + plan = self.plan if VERBOSE: - explain_plan(plan, self.total_steps, chunked=self.chunked) + self.explain(restart_chunked) - def noise_sampler(*_args): + def noise_sampler(*_args: list): return torch.randn_like(x) # Passed to the PlanItem .execute method. Most of the time, self.make_noise_sampler @@ -511,9 +670,9 @@ class KSamplerRestartWrapper: # don't all use the same noise. def get_noise_sampler(x, s_min, s_max): nonlocal seed - if not self.make_noise_sampler: + if not restart_make_noise_sampler: return noise_sampler - result = self.make_noise_sampler(x, s_min, s_max, seed) + result = restart_make_noise_sampler(x, s_min, s_max, seed) seed += 1 return result @@ -554,7 +713,7 @@ class KSamplerRestartWrapper: 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 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) @@ -566,37 +725,78 @@ class KSamplerRestartWrapper: # Execute the plan items in sequence. for pi in plan: x = pi.execute(x, do_sample, get_noise_sampler) - return x - @torch.no_grad() - def ksampler_restart_wrapper( - self, + @staticmethod + def self_test( model, - x, - sigmas, - *args, - extra_args=None, - callback=None, - disable=None, - **kwargs, - ): - plan, total_steps = build_plan( - self.real_model, - self.restart_segments, - self.restart_scheduler, - sigmas, - x.device, - ) - return self.sample_plan( - plan, - total_steps, - model, - x, - sigmas, - *args, - extra_args=extra_args, - callback=callback, - disable=disable, - **kwargs, - ) + 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 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}") + 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") diff --git a/restart_schedulers.py b/restart_schedulers.py index 012cffb..e8e78c5 100644 --- a/restart_schedulers.py +++ b/restart_schedulers.py @@ -7,7 +7,7 @@ from comfy.k_diffusion import sampling as k_diffusion_sampling # These two may be wrong for v-pred... but it seems to work? # Copied from k_diffusion def sigma_to_t(ms, sigma, quantize=True): - log_sigmas = ms.log_sigmas + log_sigmas = ms.log_sigmas.cpu() log_sigma = sigma.log() dists = log_sigma - log_sigmas[:, None] if quantize: