diff --git a/nodes.py b/nodes.py index 9e8717c..66e4a58 100644 --- a/nodes.py +++ b/nodes.py @@ -1,6 +1,12 @@ import comfy +import torch -from .restart_sampling import DEFAULT_SEGMENTS, SCHEDULER_MAPPING, restart_sampling +from . import restart_sampling as restart +from .restart_sampling import ( + DEFAULT_SEGMENTS, + SCHEDULER_MAPPING, + restart_sampling, +) def get_supported_samplers(): @@ -296,11 +302,147 @@ class KRestartSamplerCustom: ) +class RestartScheduler: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "steps": ("INT", {"default": 20, "min": 1, "max": 10000}), + "scheduler": (tuple(SCHEDULER_MAPPING.keys()),), + "segments": ( + "STRING", + {"default": DEFAULT_SEGMENTS, "multiline": False}, + ), + "restart_scheduler": (get_supported_restart_schedulers(),), + "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}), + }, + "optional": { + "sigmas_opt": ("SIGMAS",), + }, + } + + RETURN_TYPES = ("SIGMAS",) + 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 * -1 + + def go( + self, + model, + steps, + scheduler, + segments, + restart_scheduler, + 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) + + 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, + restart_scheduler, + sigmas, + "cpu", + ) + 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,) + + +class RestartSampler: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "sampler": ("SAMPLER",), + }, + } + + RETURN_TYPES = ("SAMPLER",) + FUNCTION = "go" + CATEGORY = "sampling/custom_sampling/samplers" + + def go(self, sampler): + wrapped = comfy.samplers.KSAMPLER( + lambda *args, **kwargs: self.sampler_function(sampler, *args, **kwargs), + extra_options=sampler.extra_options, + inpaint_options=sampler.inpaint_options, + ) + return (wrapped,) + + @staticmethod + @torch.no_grad() + def sampler_function(wrapped, model, x, sigmas, *args, **kwargs): + last_sigma = None + chunks = [] + while len(sigmas) > 0: + last_sigma = None + for idx in range(len(sigmas) - 1): + curr_sigma = sigmas[idx + 1] + if last_sigma is None or ( + curr_sigma.sign() == last_sigma.sign() + and curr_sigma.abs() < last_sigma.abs() + ): + last_sigma = curr_sigma + continue + break + if idx == len(sigmas) - 2: + chunks.append(sigmas) + break + chunks.append( + sigmas[: idx + 1] * -1 if sigmas[0] < 0 else sigmas[: idx + 1], + ) + sigmas = sigmas[idx + 1 :] + print("CHUNKS", chunks) + chunks = [chunk for chunk in chunks if len(chunk) > 1] + + for idx, chunk_sigmas in enumerate(chunks): + print(">>>", idx, chunk_sigmas) + if idx > 0 and chunk_sigmas[0] > chunks[idx - 1][-1]: + print("NOISE", chunk_sigmas[0], chunk_sigmas[-1]) + x += ( + torch.randn_like(x) + * (chunk_sigmas[0] ** 2 - chunk_sigmas[-1] ** 2) ** 0.5 + ) + x = wrapped.sampler_function(model, x, chunk_sigmas, *args, **kwargs) + return x + + NODE_CLASS_MAPPINGS = { "KRestartSamplerSimple": KRestartSamplerSimple, "KRestartSampler": KRestartSampler, "KRestartSamplerAdv": KRestartSamplerAdv, "KRestartSamplerCustom": KRestartSamplerCustom, + "RestartScheduler": RestartScheduler, + "RestartSampler": RestartSampler, } NODE_DISPLAY_NAME_MAPPINGS = { diff --git a/restart_sampling.py b/restart_sampling.py index 192f13f..76f6d72 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -314,6 +314,84 @@ 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, + ) + 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(plan, total_steps, chunked=True): + def pretty_sigmas(sigmas): + return ", ".join(f"{sig:.4}" for sig in sigmas.tolist()) + + print(plan) + 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 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.", + ) + + class KSamplerRestartWrapper: # Some extra explanation for a couple of these arguments: # @@ -346,81 +424,6 @@ 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) - 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, @@ -436,10 +439,16 @@ class KSamplerRestartWrapper: ksampler = self.ksampler step = 0 seed = self.seed - plan, self.total_steps = self.build_plan(sigmas, x.device) + plan, self.total_steps = build_plan( + self.real_model, + self.restart_segments, + self.restart_scheduler, + sigmas, + x.device, + ) if VERBOSE: - self.explain_plan(plan, self.total_steps) + explain_plan(plan, self.total_steps, chunked=self.chunked) def noise_sampler(*_args): return torch.randn_like(x)