From ae32ce995aa70bf9a237291ce57d3bc4079a5b53 Mon Sep 17 00:00:00 2001 From: blepping Date: Mon, 1 Apr 2024 12:39:58 -0600 Subject: [PATCH] Code formatting + clean up some lints --- nodes.py | 261 ++++++++++++++++++++++++++++++++---------- restart_sampling.py | 225 ++++++++++++++++++++++++++++-------- restart_schedulers.py | 29 ++--- 3 files changed, 393 insertions(+), 122 deletions(-) diff --git a/nodes.py b/nodes.py index 70ed2f1..9e8717c 100644 --- a/nodes.py +++ b/nodes.py @@ -1,5 +1,6 @@ import comfy -from .restart_sampling import restart_sampling, SCHEDULER_MAPPING, DEFAULT_SEGMENTS + +from .restart_sampling import DEFAULT_SEGMENTS, SCHEDULER_MAPPING, restart_sampling def get_supported_samplers(): @@ -27,129 +28,273 @@ def get_supported_restart_schedulers(): class KRestartSamplerSimple: @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): return { "required": { - "model": ("MODEL", ), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "model": ("MODEL",), + "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_name": (get_supported_samplers(), ), - "scheduler": (comfy.samplers.KSampler.SCHEDULERS, ), - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "latent_image": ("LATENT", ), - "denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "sampler_name": (get_supported_samplers(),), + "scheduler": (comfy.samplers.KSampler.SCHEDULERS,), + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "latent_image": ("LATENT",), + "denoise": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, + ), "segments": ("STRING", {"default": "default", "multiline": False}), - } + }, } RETURN_TYPES = ("LATENT",) FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments): - return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, scheduler, denoise=denoise) + def sample( + self, + model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image, + denoise, + segments, + ): + return restart_sampling( + model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image, + segments, + scheduler, + denoise=denoise, + ) class KRestartSampler: @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): return { "required": { - "model": ("MODEL", ), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "model": ("MODEL",), + "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_name": (get_supported_samplers(), ), - "scheduler": (tuple(SCHEDULER_MAPPING.keys()), ), - "positive": ("CONDITIONING", ), - "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}), - "restart_scheduler": (get_supported_restart_schedulers(), ), + "sampler_name": (get_supported_samplers(),), + "scheduler": (tuple(SCHEDULER_MAPPING.keys()),), + "positive": ("CONDITIONING",), + "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}, + ), + "restart_scheduler": (get_supported_restart_schedulers(),), "chunked_mode": ("BOOLEAN", {"default": True}), - } + }, } RETURN_TYPES = ("LATENT",) FUNCTION = "sample" CATEGORY = "sampling" - def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, segments, restart_scheduler, chunked_mode=True): - return restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, segments, restart_scheduler, denoise=denoise, chunked_mode=chunked_mode) + def sample( + self, + model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image, + denoise, + segments, + restart_scheduler, + chunked_mode=True, + ): + 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: - @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): return { "required": { "model": ("MODEL",), - "add_noise": (["enable", "disable"], ), - "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "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_name": (get_supported_samplers(), ), - "scheduler": (tuple(SCHEDULER_MAPPING.keys()), ), - "positive": ("CONDITIONING", ), - "negative": ("CONDITIONING", ), - "latent_image": ("LATENT", ), + "sampler_name": (get_supported_samplers(),), + "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(), ), + "return_with_leftover_noise": (["disable", "enable"],), + "segments": ( + "STRING", + {"default": DEFAULT_SEGMENTS, "multiline": False}, + ), + "restart_scheduler": (get_supported_restart_schedulers(),), "chunked_mode": ("BOOLEAN", {"default": True}), - } + }, } RETURN_TYPES = ("LATENT",) 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, chunked_mode=True): + 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=True, + ): 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, chunked_mode=chunked_mode) + 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: - @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): return { "required": { "model": ("MODEL",), - "add_noise": (["enable", "disable"], ), - "noise_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "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", ), + "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(), ), + "return_with_leftover_noise": (["disable", "enable"],), + "segments": ( + "STRING", + {"default": DEFAULT_SEGMENTS, "multiline": False}, + ), + "restart_scheduler": (get_supported_restart_schedulers(),), "chunked_mode": ("BOOLEAN", {"default": True}), - } + }, } - RETURN_TYPES = ("LATENT","LATENT") + 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, chunked_mode=True): + 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=True, + ): 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, 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, diff --git a/restart_sampling.py b/restart_sampling.py index 876a02c..542a1f1 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -1,14 +1,16 @@ import ast -from collections import namedtuple import os import warnings -import torch -from tqdm.auto import trange -import latent_preview +from collections import namedtuple + import comfy -from comfy.sample import sample_custom, prepare_noise +import latent_preview +import torch +from comfy.sample import prepare_noise, sample_custom from comfy.samplers import KSAMPLER, sampler_object from comfy.utils import ProgressBar +from tqdm.auto import trange + from .restart_schedulers import SCHEDULER_MAPPING VERBOSE = os.environ.get("COMFYUI_VERBOSE_RESTART_SAMPLING", "").strip() == "1" @@ -19,7 +21,7 @@ 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: restart_segments = [] - restart_segments.append({'n': n_restart, 'k': k, 't_min': t_min, 't_max': t_max}) + restart_segments.append({"n": n_restart, "k": k, "t_min": t_min, "t_max": t_max}) return restart_segments @@ -63,9 +65,9 @@ def prepare_restart_segments(restart_info, ms, sigmas): if restart_arrays is None: try: restart_arrays = ast.literal_eval(f"[{restart_info}]") - except SyntaxError as e: + except SyntaxError: print("Ill-formed restart segments") - raise e + raise restart_segments = [] for arr in restart_arrays: if len(arr) != 4: @@ -74,7 +76,13 @@ def prepare_restart_segments(restart_info, ms, sigmas): n_restart, k = int(n_restart), int(k) t_min = resolve_t_value(val_min, ms) t_max = resolve_t_value(val_max, ms) - restart_segments = add_restart_segment(restart_segments, n_restart, k, t_min, t_max) + restart_segments = add_restart_segment( + restart_segments, + n_restart, + k, + t_min, + t_max, + ) return restart_segments @@ -86,20 +94,32 @@ def round_restart_segments(ts, restart_segments): :return: dict of the form {nearest_t_min: {'n': n, 'k': k, 't_max': t_max}} """ t_min_mapping = {} - for segment in reversed(restart_segments): # Reversed to prioritize segments to the front - t_min_neighbor = min(ts, key=lambda ts: abs(ts - segment['t_min'])).item() + for segment in reversed( + restart_segments, + ): # Reversed to prioritize segments to the front + t_min_neighbor = min(ts, key=lambda ts: abs(ts - segment["t_min"])).item() if t_min_neighbor == ts[0]: warnings.warn( - f"\n[Restart Sampling] nearest neighbor of segment t_min {segment['t_min']:.4f} is equal to the first t_min in the denoise schedule {ts[0]:.4f}, ignoring segment...", stacklevel=2) + f"\n[Restart Sampling] nearest neighbor of segment t_min {segment['t_min']:.4f} is equal to the first t_min in the denoise schedule {ts[0]:.4f}, ignoring segment...", + stacklevel=2, + ) continue - if t_min_neighbor > segment['t_max']: + if t_min_neighbor > segment["t_max"]: warnings.warn( - f"\n[Restart Sampling] t_min neighbor {t_min_neighbor:.4f} is greater than t_max {segment['t_max']:.4f}, ignoring segment...", stacklevel=2) + f"\n[Restart Sampling] t_min neighbor {t_min_neighbor:.4f} is greater than t_max {segment['t_max']:.4f}, ignoring segment...", + stacklevel=2, + ) continue if t_min_neighbor in t_min_mapping: warnings.warn( - f"\n[Restart Sampling] Overwriting segment {t_min_mapping[t_min_neighbor]}, nearest neighbor of {segment['t_min']:.4f} is {t_min_neighbor:.4f}", stacklevel=2) - t_min_mapping[t_min_neighbor] = {'n': segment['n'], 'k': segment['k'], 't_max': segment['t_max']} + f"\n[Restart Sampling] Overwriting segment {t_min_mapping[t_min_neighbor]}, nearest neighbor of {segment['t_min']:.4f} is {t_min_neighbor:.4f}", + stacklevel=2, + ) + t_min_mapping[t_min_neighbor] = { + "n": segment["n"], + "k": segment["k"], + "t_max": segment["t_max"], + } return t_min_mapping @@ -110,11 +130,31 @@ def calc_sigmas(scheduler, n, sigma_min, sigma_max, model, device): def calc_restart_steps(restart_segments): restart_steps = 0 for segment in restart_segments.values(): - restart_steps += (segment['n'] - 1) * segment['k'] + restart_steps += (segment["n"] - 1) * segment["k"] 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=True, sigmas=None): +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, + sigmas=None, +): if isinstance(sampler, str): sampler = sampler_object(sampler) @@ -123,11 +163,17 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega 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) + 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, + 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) @@ -135,26 +181,45 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega start_step, last_step = step_range if last_step < (len(sigmas) - 1): - sigmas = sigmas[:last_step + 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):] + sigmas = sigmas[-(steps + 1) :] - restart_segments = prepare_restart_segments(restart_info, real_model.model_sampling, sigmas) + 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) + sampler_wrapper = KSamplerRestartWrapper( + sampler, + real_model, + restart_scheduler, + restart_segments, + seed, + custom_noise, + chunked=chunked_mode, + ) latent = latent_image latent_image = latent["samples"] if disable_noise: - torch.manual_seed(seed) # workaround for https://github.com/comfyanonymous/ComfyUI/issues/2833 - noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu") + torch.manual_seed( + seed, + ) # workaround for https://github.com/comfyanonymous/ComfyUI/issues/2833 + noise = torch.zeros( + latent_image.size(), + dtype=latent_image.dtype, + layout=latent_image.layout, + device="cpu", + ) else: - batch_inds = latent["batch_index"] if "batch_index" in latent else None + batch_inds = latent.get("batch_index", None) noise = prepare_noise(latent_image, seed, batch_inds) noise_mask = None @@ -167,11 +232,11 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED sampler = KSAMPLER( - sampler_wrapper.ksampler_restart_wrapper, extra_options=sampler.extra_options | {}, + sampler_wrapper.ksampler_restart_wrapper, + extra_options=sampler.extra_options | {}, inpaint_options=sampler.inpaint_options | {}, ) - # Add the additional steps to the progress bar pbar_update_absolute = ProgressBar.update_absolute @@ -182,8 +247,19 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega try: samples = sample_custom( - model, noise, cfg, sampler, sigmas, positive, negative, latent_image, - noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=seed) + model, + noise, + cfg, + sampler, + sigmas, + positive, + negative, + latent_image, + noise_mask=noise_mask, + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) finally: ProgressBar.update_absolute = pbar_update_absolute @@ -202,7 +278,13 @@ def restart_sampling(model, seed, steps, cfg, sampler, scheduler, positive, nega return (out, out_denoised) -class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart_sigmas"], defaults=[None, 0, 0., 0., None])): +class PlanItem( + namedtuple( + "PlanItem", + ["sigmas", "k", "s_min", "s_max", "restart_sigmas"], + defaults=[None, 0, 0.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. @@ -224,7 +306,10 @@ class PlanItem(namedtuple("PlanItem", ["sigmas", "k", "s_min", "s_max", "restart 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 + x += ( + noise_sampler(self.restart_sigmas[0], self.restart_sigmas[-1]) + * (self.s_max**2 - self.s_min**2) ** 0.5 + ) x = sample(x, self.restart_sigmas, kidx) return x @@ -242,7 +327,16 @@ 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): + 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 @@ -269,10 +363,18 @@ class KSamplerRestartWrapper: 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])) + 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. @@ -284,9 +386,11 @@ class KSamplerRestartWrapper: 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. @@ -295,13 +399,15 @@ class KSamplerRestartWrapper: 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): + 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)}") + print( + f"[{rlabel}] Step {step:>3}..{step+chunk_size:<3}: {pretty_sigmas(sigs)}", + ) step += chunk_size return x @@ -311,11 +417,22 @@ class KSamplerRestartWrapper: 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.") - + print( + "** Plan legend: [Rn] - steps for restart #n, normal sampling steps otherwise. Ranges are inclusive.", + ) @torch.no_grad() - def ksampler_restart_wrapper(self, model, x, sigmas, *args, extra_args=None, callback=None, disable=None, **kwargs): + def ksampler_restart_wrapper( + self, + model, + x, + sigmas, + *args, + extra_args=None, + callback=None, + disable=None, + **kwargs, + ): ksampler = self.ksampler step = 0 seed = self.seed @@ -340,6 +457,7 @@ class KSamplerRestartWrapper: return result with trange(self.total_steps, disable=disable) as pbar: + def callback_wrapper(x): nonlocal step step += 1 @@ -351,10 +469,17 @@ class KSamplerRestartWrapper: # 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) + model, + x, + sigs, + *args, + extra_args=extra_args, + callback=callback_wrapper, + disable=True, + **kwargs, + ) - def do_sample(x, sigs, kidx=-1): + 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: @@ -362,8 +487,8 @@ class KSamplerRestartWrapper: # 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]) + for i in range(len(sigs) - 1): + x = sampler_function(x, sigs[i : i + 2]) return x # Execute the plan items in sequence. diff --git a/restart_schedulers.py b/restart_schedulers.py index 82d9b55..793116c 100644 --- a/restart_schedulers.py +++ b/restart_schedulers.py @@ -1,7 +1,9 @@ import torch from comfy.k_diffusion import sampling as k_diffusion_sampling + # from comfy.samplers import normal_scheduler + # 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): @@ -18,6 +20,7 @@ def sigma_to_t(ms, sigma, quantize=True): t = (1 - w) * low_idx + w * high_idx return t.view(sigma.shape) + # Copied from k_diffusion def t_to_sigma(ms, t): t = t.float() @@ -25,15 +28,16 @@ def t_to_sigma(ms, t): log_sigma = (1 - w) * ms.log_sigmas[low_idx] + w * ms.log_sigmas[high_idx] return log_sigma.exp() -def get_sigmas_karras(model, n, s_min, s_max, device): + +def get_sigmas_karras(_model, n, s_min, s_max, device): return k_diffusion_sampling.get_sigmas_karras(n, s_min, s_max, device=device) -def get_sigmas_exponential(model, n, s_min, s_max, device): +def get_sigmas_exponential(_model, n, s_min, s_max, device): return k_diffusion_sampling.get_sigmas_exponential(n, s_min, s_max, device=device) -def normal_scheduler(model, steps, s_min, s_max, sgm=False, floor=False): +def normal_scheduler(model, steps, s_min, s_max, sgm=False): """ Pulled from comfy.samplers.normal_scheduler """ @@ -46,7 +50,7 @@ def normal_scheduler(model, steps, s_min, s_max, sgm=False, floor=False): else: timesteps = torch.linspace(start, end, steps) - sigs = tuple(ms.sigma(timesteps[x]) for x in range(len(timesteps))) + (0.0,) + sigs = (*(ms.sigma(timesteps[x]) for x in range(len(timesteps))), 0.0) return torch.FloatTensor(sigs) @@ -60,10 +64,12 @@ def get_sigmas_simple(model, n, s_min, s_max, device): max_idx = torch.argmin(torch.abs(ms.sigmas - s_max)) sigmas_slice = ms.sigmas[min_idx:max_idx] ss = len(sigmas_slice) / n - sigs = [float(s_max)] - for x in range(1, n - 1): - sigs += [float(sigmas_slice[-(1 + int(x * ss))])] - sigs += [float(s_min), 0.0] + sigs = ( + float(s_max), + *(float(sigmas_slice[-(1 + int(x * ss))]) for x in range(1, n - 1)), + float(s_min), + 0.0, + ) return torch.tensor(sigs, device=device) @@ -71,12 +77,7 @@ def get_sigmas_ddim_uniform(model, n, s_min, s_max, device): ms = model.model_sampling t_min, t_max = sigma_to_t(ms, torch.tensor([s_min, s_max], device=device)) ddim_timesteps = torch.linspace(t_max, t_min, n, dtype=torch.int16, device=device) - sigs = [] - for ts in ddim_timesteps: - if ts > 999: - ts = 999 - sigs.append(t_to_sigma(ms, ts)) - sigs += [0.0] + sigs = (*(t_to_sigma(ms, min(ts, 999)) for ts in ddim_timesteps), 0.0) return torch.tensor(sigs, device=device)