From 9258cb2d90fb2fd1bc5462ed9dba05f12ffa4f99 Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 9 Jun 2024 21:57:44 -0600 Subject: [PATCH] Stage 1 --- __init__.py | 3 + py/nodes.py | 273 ++++++++++++++++++++++-------- py/sampling.py | 60 +++---- py/substep_merging.py | 370 +++++++++++++++++++++++++++-------------- py/substep_samplers.py | 317 +++++++++++++++++++++++++++++------ py/substep_sampling.py | 235 ++++++++++++++++++++++++-- py/utils.py | 40 ++++- 7 files changed, 997 insertions(+), 301 deletions(-) diff --git a/__init__.py b/__init__.py index f7e3aba..4122a89 100644 --- a/__init__.py +++ b/__init__.py @@ -4,5 +4,8 @@ from .py import nodes NODE_CLASS_MAPPINGS = { "ComposableSampler": nodes.ComposableSampler, "ComposableStepSampler": nodes.ComposableStepSampler, + "SubstepsGroup": nodes.SubstepsGroup, + "CSamplerParam": nodes.CSamplerParam, + "CSamplerParamMulti": nodes.CSamplerParamMulti, } __all__ = ["NODE_CLASS_MAPPINGS"] diff --git a/py/nodes.py b/py/nodes.py index 3b5b764..dd01878 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -1,10 +1,17 @@ -from .sampling import composable_sampler, STEP_SAMPLERS -from .substep_sampling import StepSamplerChain +from .sampling import composable_sampler +from .substep_sampling import StepSamplerChain, StepSamplerGroups, ParamGroup +from .substep_samplers import STEP_SAMPLERS from .substep_merging import MERGE_SUBSTEPS_CLASSES import comfy import yaml +DEFAULT_YAML_PARAMS = """\ +# Enter parameters here in JSON or YAML format +s_noise: 1.0 +eta: 1.0 +""" + class ComposableSampler: RETURN_TYPES = ("SAMPLER",) @@ -16,34 +23,17 @@ class ComposableSampler: def INPUT_TYPES(cls): return { "required": { - "s_noise": ( - "FLOAT", - { - "default": 1.0, - "min": -100.0, - "max": 100.0, - "step": 0.01, - "round": False, - }, - ), - "eta": ( - "FLOAT", - { - "default": 1.0, - "min": -100.0, - "max": 100.0, - "step": 0.01, - "round": False, - }, - ), - "merge_method": (tuple(MERGE_SUBSTEPS_CLASSES.keys()),), - "step_sampler_chain": ("STEP_SAMPLER_CHAIN",), + "step_sampler_groups": ("STEP_SAMPLER_GROUPS",), }, "optional": { - "merge_sampler_opt": ("STEP_SAMPLER_CHAIN",), + "csampler_params_opt": ("CSAMPLER_PARAMS",), "parameters": ( "STRING", - {"default": "", "multiline": True, "dynamicPrompts": False}, + { + "default": DEFAULT_YAML_PARAMS, + "multiline": True, + "dynamicPrompts": False, + }, ), }, } @@ -51,30 +41,21 @@ class ComposableSampler: def go( self, *, - s_noise, - eta, - merge_method, - step_sampler_chain, - merge_sampler_opt=None, + step_sampler_groups, + csampler_params_opt=None, parameters="", ): - if merge_sampler_opt is not None: - merge_sampler = merge_sampler_opt.items[0] - else: - merge_sampler = ComposableStepSampler().go(step_method="euler")[0].items[0] - options = { - "s_noise": s_noise, - "eta": eta, - "merge_method": merge_method, - "merge_sampler": merge_sampler, - } + options = {} parameters = parameters.strip() if parameters: extra_params = yaml.safe_load(parameters) - if not isinstance(extra_params, dict): - raise ValueError("Parameters must be a JSON or YAML object") - options |= extra_params - options["chain"] = step_sampler_chain.clone() + if extra_params is not None: + if not isinstance(extra_params, dict): + raise ValueError("Parameters must be a JSON or YAML object") + options |= extra_params + if csampler_params_opt is not None: + options |= csampler_params_opt.items + options["_groups"] = step_sampler_groups.clone() return ( comfy.samplers.KSAMPLER( composable_sampler, @@ -83,6 +64,78 @@ class ComposableSampler: ) +class SubstepsGroup: + RETURN_TYPES = ("STEP_SAMPLER_GROUPS",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "go" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "merge_method": (tuple(MERGE_SUBSTEPS_CLASSES.keys()),), + "time_mode": (("step", "step_pct", "sigma"),), + "time_start": ( + "FLOAT", + {"default": 0, "min": 0.0, "step": 0.1, "round'": False}, + ), + "time_end": ( + "FLOAT", + {"default": 999, "min": 0.0, "step": 0.1, "round'": False}, + ), + "step_sampler_chain": ("STEP_SAMPLER_CHAIN",), + }, + "optional": { + "step_sampler_groups_opt": ("STEP_SAMPLER_GROUPS",), + "csampler_params_opt": ("CSAMPLER_PARAMS",), + "parameters": ( + "STRING", + { + "default": DEFAULT_YAML_PARAMS, + "multiline": True, + "dynamicPrompts": False, + }, + ), + }, + } + + def go( + self, + *, + merge_method, + time_mode, + time_start, + time_end, + step_sampler_chain, + step_sampler_group_opt=None, + csampler_params_opt=None, + parameters="", + ): + group = ( + StepSamplerGroups() + if step_sampler_group_opt is None + else step_sampler_group_opt + ) + chain = step_sampler_chain.clone() + chain.merge_method = merge_method + chain.time_mode = time_mode + chain.time_start, chain.time_end = time_start, time_end + options = {} + parameters = parameters.strip() + if parameters: + extra_params = yaml.safe_load(parameters) + if extra_params is not None: + if not isinstance(extra_params, dict): + raise ValueError("Parameters must be a JSON or YAML object") + options |= extra_params + if csampler_params_opt is not None: + options |= csampler_params_opt.items + chain.options |= options + group.append(chain) + return (group,) + + class ComposableStepSampler: RETURN_TYPES = ("STEP_SAMPLER_CHAIN",) CATEGORY = "sampling/custom_sampling/samplers" @@ -93,40 +146,31 @@ class ComposableStepSampler: def INPUT_TYPES(cls): return { "required": { - "s_noise": ( - "FLOAT", - { - "default": 1.0, - "min": -100.0, - "max": 100.0, - "step": 0.01, - "round": False, - }, - ), - "eta": ( - "FLOAT", - { - "default": 1.0, - "min": -100.0, - "max": 100.0, - "step": 0.01, - "round": False, - }, - ), "substeps": ("INT", {"default": 1, "min": 1, "max": 1000}), "step_method": (tuple(STEP_SAMPLERS.keys()),), }, "optional": { "step_sampler_opt": ("STEP_SAMPLER_CHAIN",), - "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), + "csampler_params_opt": ("CSAMPLER_PARAMS",), "parameters": ( "STRING", - {"default": "", "multiline": True, "dynamicPrompts": False}, + { + "default": DEFAULT_YAML_PARAMS, + "multiline": True, + "dynamicPrompts": False, + }, ), }, } - def go(self, *, parameters="", step_sampler_opt=None, **kwargs): + def go( + self, + *, + parameters="", + step_sampler_opt=None, + csampler_params_opt=None, + **kwargs, + ): if step_sampler_opt is not None: chain = step_sampler_opt.clone() else: @@ -134,11 +178,94 @@ class ComposableStepSampler: parameters = parameters.strip() if parameters: extra_params = yaml.safe_load(parameters) - if not isinstance(extra_params, dict): - raise ValueError("Parameters must be a JSON or YAML object") - kwargs |= extra_params - chain.items.append(kwargs) + if extra_params is not None: + if not isinstance(extra_params, dict): + raise ValueError("Parameters must be a JSON or YAML object") + kwargs |= extra_params + if csampler_params_opt is not None: + kwargs |= csampler_params_opt.items + chain.append(kwargs) return (chain,) -__all__ = ("ComposableStepSampler", "ComposableSampler") +class Wildcard(str): + __slots__ = () + + def __ne__(self, _unused): + return False + + +class CSamplerParam: + RETURN_TYPES = ("CSAMPLER_PARAMS",) + CATEGORY = "sampling/custom_sampling/samplers" + FUNCTION = "go" + + WC = Wildcard("*") + + CPARAM_TYPES = { + "custom_noise": lambda v: hasattr(v, "make_noise_sampler"), + "merge_sampler": lambda v: isinstance(v, StepSamplerChain), + } + + @classmethod + def INPUT_TYPES(cls): + return { + "required": {"key": (tuple(cls.CPARAM_TYPES.keys()),), "value": (cls.WC,)}, + "optional": {"csampler_params_opt": ("CSAMPLER_PARAMS",)}, + } + + def go(self, *, key, value, csampler_params_opt=None): + if not self.CPARAM_TYPES[key](value): + raise ValueError(f"CSamplerParam: Bad value type for key {key}") + params = ( + ParamGroup(items={}) + if csampler_params_opt is None + else csampler_params_opt.clone() + ) + params[key] = value + return (params,) + + +class CSamplerParamMulti: + RETURN_TYPES = ("CSAMPLER_PARAMS",) + CATEGORY = "sampling/custom_sampling/samplers" + FUNCTION = "go" + + PARAM_COUNT = 5 + + @classmethod + def INPUT_TYPES(cls): + param_keys = (("", *CSamplerParam.CPARAM_TYPES.keys()),) + return { + "required": { + f"key_{idx}": param_keys for idx in range(1, cls.PARAM_COUNT + 1) + }, + "optional": {"csampler_params_opt": ("CSAMPLER_PARAMS",)} + | { + f"value_opt_{idx}": (CSamplerParam.WC,) + for idx in range(1, cls.PARAM_COUNT + 1) + }, + } + + def go(self, *, csampler_params_opt=None, **kwargs): + params = ( + ParamGroup(items={}) + if csampler_params_opt is None + else csampler_params_opt.clone() + ) + for idx in range(1, self.PARAM_COUNT + 1): + key, value = kwargs.get(f"key_{idx}"), kwargs.get(f"value_opt_{idx}") + if not key or value is None: + continue + if not CSamplerParam.CPARAM_TYPES[key](value): + raise ValueError(f"CSamplerParamGroup: Bad value type for key {key}") + params[key] = value + return (params,) + + +__all__ = ( + "ComposableStepSampler", + "ComposableSampler", + "CSamplerParam", + "CSamplerParamMulti", +) diff --git a/py/sampling.py b/py/sampling.py index d373c20..2af4893 100644 --- a/py/sampling.py +++ b/py/sampling.py @@ -2,8 +2,7 @@ import torch from tqdm.auto import trange -from .substep_samplers import STEP_SAMPLERS -from .substep_sampling import SamplerState, History, ModelCallCache +from .substep_sampling import SamplerState, History, ModelCallCache, NoiseSamplerCache from .substep_merging import MERGE_SUBSTEPS_CLASSES @@ -29,35 +28,6 @@ def composable_sampler( def noise_sampler(_s, _sn): return torch.randn_like(x) - samplers = [] - substeps = 0 - for sitem in copts["chain"].items: - custom_noise = sitem.get("custom_noise_opt") - if custom_noise is None: - curr_ns = noise_sampler - else: - curr_ns = custom_noise.make_noise_sampler( - x, sigmas[-1], sigmas[0], normalized=True - ) - ssampler = STEP_SAMPLERS[sitem["step_method"]](noise_sampler=curr_ns, **sitem) - samplers.append(ssampler) - # samplers += (ssampler,) * sitem["substeps"] - substeps += ssampler.substeps - msitem = copts["merge_sampler"] - if copts["merge_method"] in ("sample", "sample_uncached"): - custom_noise = msitem.get("custom_noise_opt") - if custom_noise is None: - curr_ns = noise_sampler - else: - curr_ns = custom_noise.make_noise_sampler( - x, sigmas[-1], sigmas[0], normalized=True - ) - merge_sampler = STEP_SAMPLERS[msitem["step_method"]]( - noise_sampler=curr_ns, **msitem - ) - pass - else: - merge_sampler = None ss = SamplerState( ModelCallCache( model, @@ -79,14 +49,30 @@ def composable_sampler( s_noise=s_noise if s_noise != 1.0 else copts["s_noise"], reta=copts.get("reta", 1.0), ) - merge_sampler = MERGE_SUBSTEPS_CLASSES[copts["merge_method"]]( - ss, - samplers, - **(copts | {"merge_sampler": merge_sampler}), + groups = copts["_groups"] + merge_samplers = tuple( + MERGE_SUBSTEPS_CLASSES[g.merge_method](ss, g.items, **g.options) + for g in groups.items ) - for idx in trange(len(sigmas) - 1, disable=disable): - print(f"STEP {idx+1}") + nsc = NoiseSamplerCache( + x, + extra_args.get("seed", 42), + sigmas[-1], + sigmas[0], + **copts.get("noise", {}), + ) + ss.noise = nsc + step_count = len(sigmas) - 1 + for idx in trange(step_count, disable=disable): + print(f"STEP {idx + 1}") ss.update(idx) ss.model.reset_cache() + nsc.update_x(x) + ms_idx = groups.find_match(ss.sigma, idx, step_count) + if ms_idx is None: + raise RuntimeError(f"No matching sampler group for step {idx + 1}") + merge_sampler = merge_samplers[ms_idx] x = merge_sampler.step(x) + if (idx + 1) % nsc.cache_reset_interval == 0: + nsc.reset_cache() return x diff --git a/py/substep_merging.py b/py/substep_merging.py index 8230dc4..d51d06c 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -1,27 +1,83 @@ import torch +import contextlib -from .utils import scale_noise, find_first_unsorted +from .utils import find_first_unsorted +from .substep_samplers import STEP_SAMPLERS from .substep_sampling import History class MergeSubstepsSampler: - def __init__(self, ss, samplers, **_kwargs): + def __init__(self, ss, sitems, **kwargs): + samplers = tuple( + STEP_SAMPLERS[sitem["step_method"]](**sitem) for sitem in sitems + ) self.ss = ss self.samplers = samplers self.substeps = sum(sampler.substeps for sampler in samplers) + self.options = kwargs def step(self, x): raise NotImplementedError - def merge_steps(self, _x, result): + def substep(self, x, sampler, ss=None): + if ss is None: + ss = self.ss + sg = sampler.step(x, ss) + next_x = None + with contextlib.suppress(StopIteration): + while True: + sr = sg.send(next_x) + next_x = sr.x + yield sr + + def simple_substep(self, x, sampler, ss=None): + for sr in self.substep(x, sampler, ss=ss): + if not sr.final: + sr.noise_x() + return sr + + def merge_steps(self, x, result=None, *, noise=None, ss=None): + if result is None: + result = x + ss = ss if ss is not None else self.ss + ss.dhist.push(ss.denoised) + ss.denoised = None + ss.callback(result) + if noise is not None: + result = result + noise + ss.xhist.push(result) return result + def step_max_noise_samples(self): + return sum( + (1 + sampler.self_noise) * sampler.substeps for sampler in self.samplers + ) + + +class SimpleSubstepsSampler(MergeSubstepsSampler): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if not len(self.samplers): + raise ValueError("Missing sampler") + + def step_max_noise_samples(self): + return 1 + self.samplers[0].self_noise + + def step(self, x): + ss, ssampler = self.ss, self.samplers[0] + custom_noise = ssampler.options.get( + "custom_noise", self.options.get("custom_noise") + ) + noise_sampler = ss.noise.make_caching_noise_sampler( + custom_noise, 1, ss.sigma, ss.sigma_next + ) + ssampler.noise_sampler = noise_sampler + ss.denoised = ss.model(x, ss.sigma) + sr = self.simple_substep(x, ssampler) + return self.merge_steps(sr.x, noise=sr.get_noise()) + class NormalMergeSubstepsSampler(MergeSubstepsSampler): - def __init__(self, ss, samplers, **kwargs): - super().__init__(ss, samplers, **kwargs) - self.ss = ss - def step(self, x): ss = self.ss substeps = self.substeps @@ -29,38 +85,49 @@ class NormalMergeSubstepsSampler(MergeSubstepsSampler): z_avg = torch.zeros_like(x) noise = torch.zeros_like(x) noise_total = 0.0 - for idx, ssampler in enumerate( - sampler for sampler in self.samplers for _ in range(sampler.substeps) - ): - print(f" SUBSTEP {idx+1}: {ssampler.name}") - ss.denoised = ss.model(x, ss.sigma) - z_k, noise_strength = ssampler.step(x, ss) - z_avg += renoise_weight * z_k - noise_strength *= ssampler.s_noise - if ss.sigma_next == 0 or noise_strength == 0: - continue - noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next) - x = z_k - if idx != substeps - 1: - x += noise_curr * noise_strength - noise_total += noise_strength.item() * renoise_weight - noise += noise_curr * noise_strength - ss.dhist.push(ss.denoised) - ss.denoised = None - x = self.merge_steps(x, z_avg) - if ss.sigma_next != 0 and noise_total != 0: - x += scale_noise(noise, noise_total * ss.s_noise) - ss.xhist.push(x) - ss.callback(x) - return x + substep = 0 + for ssampler in self.samplers: + custom_noise = ssampler.options.get( + "custom_noise", self.options.get("custom_noise") + ) + noise_sampler = ss.noise.make_caching_noise_sampler( + custom_noise, ssampler.max_noise_samples(), ss.sigma, ss.sigma_next + ) + ssampler.noise_sampler = noise_sampler + for subidx in range(ssampler.substeps): + print(f" SUBSTEP {substep + 1}: {ssampler.name}") + self.ss.denoised = ss.denoised = ss.model(x, ss.sigma) + sr = self.simple_substep(x, ssampler) + x = sr.x + z_avg += renoise_weight * x + noise_strength = sr.noise_scale + substep += 1 + if ss.sigma_next == 0 or noise_strength == 0: + continue + noise_curr = sr.get_noise() + if substep < substeps: + x = x + noise_curr + noise_total += noise_strength.item() * renoise_weight + noise += noise_curr + return self.merge_steps( + x, + z_avg, + noise=None + if not noise_total + else ss.noise.scale_noise(noise, noise_total * ss.s_noise, normalized=True), + ) class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler): - def __init__(self, ss, samplers, *, avgmerge_stretch=0.4, **kwargs): - super().__init__(ss, samplers, **kwargs) - self.ss = ss + def __init__(self, ss, sitems, *, avgmerge_stretch=0.4, **kwargs): + super().__init__(ss, sitems, **kwargs) self.stretch = avgmerge_stretch + def step_max_noise_samples(self): + return sum( + 1 + (2 + sampler.self_noise) * sampler.substeps for sampler in self.samplers + ) + def step(self, x): ss = orig_ss = self.ss substeps = self.substeps @@ -71,44 +138,63 @@ class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler): sig_adj = ss.sigma + stretch ss = self.ss.clone_edit(sigma=sig_adj) orig_x = x - x = x + ss.noise_sampler(orig_ss.sigma, ss.sigma_next) * stretch * ss.s_noise - ss.denoised = ss.model(x, sig_adj) + stretch_strength = stretch * ss.s_noise + if stretch_strength != 0: + noise_sampler = ss.noise.make_caching_noise_sampler( + self.options.get("custom_noise"), 1, orig_ss.sigma, ss.sigma_next + ) + x = x + ( + noise_sampler(orig_ss.sigma, ss.sigma_next).mul_(stretch * ss.s_noise) + ) + self.ss.denoised = ss.denoised = ss.model(x, sig_adj) noise_total = 0.0 - step = 0 + substep = 0 for idx, ssampler in enumerate(self.samplers): print( - f" SUBSTEP {step+1} .. {step+ssampler.substeps}: {ssampler.name}, stretch={stretch}" + f" SUBSTEP {substep + 1} .. {substep + ssampler.substeps}: {ssampler.name}, stretch={stretch}" ) + custom_noise = ssampler.options.get( + "custom_noise", self.options.get("custom_noise") + ) + noise_sampler = ss.noise.make_caching_noise_sampler( + custom_noise, + ssampler.substeps + + (0 if ss.sigma_next == 0 else ssampler.max_noise_samples()), + ss.sigma, + ss.sigma_next, + ) + ssampler.noise_sampler = noise_sampler for sidx in range(ssampler.substeps): - curr_x = orig_x + scale_noise( - ssampler.noise_sampler(sig_adj, ss.sigma_next), stretch - ) - z_k, noise_strength = ssampler.step(curr_x, ss) - z_avg += renoise_weight * z_k - if ss.sigma_next == 0: + curr_x = orig_x + noise_sampler(sig_adj, ss.sigma_next).mul_(stretch) + sr = self.simple_substep(curr_x, ssampler, ss=ss) + z_avg += renoise_weight * sr.x + noise_strength = sr.noise_scale + if ss.sigma_next == 0 or noise_strength == 0: continue - noise_strength *= ssampler.s_noise - if noise_strength == 0: - continue - noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next) + noise_curr = sr.get_noise() noise_total += noise_strength.item() * renoise_weight - noise += noise_curr * noise_strength - step += ssampler.substeps - ss.dhist.push(ss.denoised) - ss.denoised = None - x = self.merge_steps(x, z_avg) - if ss.sigma_next != 0 and noise_total != 0: - x += scale_noise(noise, noise_total * ss.s_noise) - ss.xhist.push(x) - ss.callback(x) - return x + noise += noise_curr + substep += ssampler.substeps + return self.merge_steps( + x, + z_avg, + noise=None + if not noise_total + else ss.noise.scale_noise(noise, noise_total * ss.s_noise, normalized=True), + ss=ss, + ) class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): cache_model = True - def __init__(self, ss, samplers, *, merge_sampler, **kwargs): - super().__init__(ss, samplers, **kwargs) + def __init__(self, ss, sitems, *, merge_sampler=None, **kwargs): + super().__init__(ss, sitems, **kwargs) + if merge_sampler is None: + merge_sampler = STEP_SAMPLERS["euler"](step_method="euler") + else: + msitem = merge_sampler.items[0] + merge_sampler = STEP_SAMPLERS[msitem["step_method"]](**msitem) self.merge_sampler = merge_sampler self.merge_ss = None @@ -125,41 +211,40 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): step = 0 for idx, ssampler in enumerate(self.samplers): print( - f" SUBSTEP {step+1} .. {step+ssampler.substeps}: {ssampler.name}, stretch={stretch}" + f" SUBSTEP {step + 1} .. {step + ssampler.substeps}: {ssampler.name}, stretch={stretch}" ) - if idx == 0 or not self.cache_model: - ss.denoised = ss.model( - curr_x, - # + ss.noise_sampler(sig_adj.sigma, ss.sigma_next) * stretch * ss.s_noise, - sig_adj, - ) + custom_noise = ssampler.options.get( + "custom_noise", self.options.get("custom_noise") + ) + noise_sampler = ss.noise.make_caching_noise_sampler( + custom_noise, + ssampler.max_noise_samples() + ssampler.substeps, + ss.sigma, + ss.sigma_next, + ) + ssampler.noise_sampler = noise_sampler for sidx in range(ssampler.substeps): - curr_x = ( - x - + ssampler.noise_sampler(sig_adj, ss.sigma_next) - * ssampler.s_noise - * stretch - ) - z_k, noise_strength = ssampler.step(curr_x, ss) - z_avg += renoise_weight * z_k - curr_x = z_k - if noise_strength == 0 or ss.sigma_next == 0: - continue - curr_x += ( - ssampler.noise_sampler(ss.sigma, ss.sigma_next) - * ssampler.s_noise - * noise_strength + if idx + sidx == 0 or not self.cache_model: + self.ss.denoised = ss.denoised = ss.model( + curr_x, + ss.sigma, + # + ss.noise_sampler(sig_adj.sigma, ss.sigma_next) * stretch * ss.s_noise, + # sig_adj, + ) + curr_x = x + noise_sampler(sig_adj, ss.sigma_next).mul_( + ssampler.s_noise * stretch ) + sr = self.simple_substep(curr_x, ssampler, ss=ss) + z_avg += renoise_weight * sr.x + curr_x = sr.noise_x(sr.x) step += ssampler.substeps - ss.dhist.push(ss.denoised) - ss.denoised = None - x = self.merge_steps(curr_x, z_avg) - ss.xhist.push(x) - ss.callback(x) - return x + return self.merge_steps(curr_x, z_avg) def merge_steps(self, x, result): - self.ss.model.reset_cache() + ss = self.ss + ss.dhist.push(ss.denoised) + ss.denoised = None + ss.model.reset_cache() msampler = self.merge_sampler if self.merge_ss is None: merge_ss = self.merge_ss = self.ss.clone_edit( @@ -168,26 +253,27 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): xhist=History(x, 2), s_noise=msampler.s_noise, eta=msampler.eta, - # model_call_cache=None, ) else: merge_ss = self.merge_ss merge_ss.denoised = result merge_ss.update(self.ss.idx) final = merge_ss.sigma_next == 0 - merged, noise_strength = msampler.step(x, merge_ss) - if not final: - ss = self.ss - merged = ( - merged - + msampler.noise_sampler(ss.sigma, ss.sigma_next) - * msampler.s_noise - * ss.sigma_up - ) + noise_sampler = merge_ss.noise.make_caching_noise_sampler( + msampler.options.get("custom_noise", self.options.get("custom_noise")), + msampler.max_noise_samples() + int(not final), + merge_ss.sigma, + merge_ss.sigma_next, + ) + msampler.noise_sampler = noise_sampler + sr = self.simple_substep(x, msampler, ss=merge_ss) + self.ss.callback(sr.x) + sr.noise_x() merge_ss.dhist.push(result) - merge_ss.xhist.push(merged) + merge_ss.xhist.push(sr.x) merge_ss.denoised = None - return merged + ss.xhist.push(sr.x) + return sr.x class SampleUncachedMergeSubstepsSampler(SampleMergeSubstepsSampler): @@ -195,8 +281,8 @@ class SampleUncachedMergeSubstepsSampler(SampleMergeSubstepsSampler): class DivideMergeSubstepsSampler(MergeSubstepsSampler): - def __init__(self, ss, samplers, *, schedule_multiplier=4, **kwargs): - super().__init__(ss, samplers, **kwargs) + def __init__(self, ss, sitems, *, schedule_multiplier=4, **kwargs): + super().__init__(ss, sitems, **kwargs) self.schedule_multiplier = schedule_multiplier def make_schedule(self, ss): @@ -228,28 +314,67 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler): subss = self.ss.clone_edit(idx=0, sigmas=self.make_schedule(ss)) subss.main_idx = ss.idx subss.main_sigmas = ss.sigmas - - for idx, ssampler in enumerate( - sampler for sampler in self.samplers for _ in range(sampler.substeps) - ): - print(f" SUBSTEP {idx+1}: {ssampler.name}") - subss.update(idx) - subss.denoised = subss.model(x, subss.sigma) - x, noise_strength = ssampler.step(x, subss) - if noise_strength == 0 or subss.sigma_next == 0: - continue - x = ( - x - + ssampler.noise_sampler(subss.sigma, subss.sigma_next) - * ssampler.s_noise - * noise_strength + substep = 0 + for ssampler in self.samplers: + custom_noise = ssampler.options.get( + "custom_noise", self.options.get("custom_noise") ) - subss.xhist.push(x) - subss.dhist.push(subss.denoised) - subss.denoised = None - ss.callback(x) + noise_sampler = ss.noise.make_caching_noise_sampler( + custom_noise, ssampler.max_noise_samples(), ss.sigma, ss.sigma_next + ) + ssampler.noise_sampler = noise_sampler + for subidx in range(ssampler.substeps): + print(f" SUBSTEP {substep + 1}: {ssampler.name}") + subss.update(substep) + subss.denoised = subss.model(x, subss.sigma) + sr = self.simple_substep(x, ssampler, ss=subss) + x = sr.x + noise_strength = sr.noise_scale + subss.dhist.push(subss.denoised) + subss.denoised = None + if substep == self.substeps - 1: + ss.callback(x) + if noise_strength != 0 and subss.sigma_next != 0: + x = sr.noise_x() + subss.xhist.push(x) + substep += 1 + # subss.xhist.push(x) + # subss.dhist.push(subss.denoised) + # subss.denoised = None return x + # def step(self, x): + # ss = self.ss + # # print("SUBSIGMAS", subsigmas) + # subss = self.ss.clone_edit(idx=0, sigmas=self.make_schedule(ss)) + # subss.main_idx = ss.idx + # subss.main_sigmas = ss.sigmas + # substep = 0 + # for ssampler in self.samplers: + # custom_noise = ssampler.options.get( + # "custom_noise", self.options.get("custom_noise") + # ) + # noise_sampler = ss.noise.make_caching_noise_sampler( + # custom_noise, ssampler.max_noise_samples(), ss.sigma, ss.sigma_next + # ) + # ssampler.noise_sampler = noise_sampler + # for subidx in range(ssampler.substeps): + # print(f" SUBSTEP {substep + 1}: {ssampler.name}") + # subss.update(substep) + # subss.denoised = subss.model(x, subss.sigma) + # sr = self.simple_substep(x, ssampler, ss=subss) + # x = sr.x + # noise_strength = sr.noise_scale + # subss.dhist.push(subss.denoised) + # subss.denoised = None + # if substep == self.substeps - 1: + # ss.callback(x) + # if noise_strength != 0 and subss.sigma_next != 0: + # x = sr.noise_x() + # subss.xhist.push(x) + # substep += 1 + # return x + MERGE_SUBSTEPS_CLASSES = { "normal": NormalMergeSubstepsSampler, @@ -257,4 +382,5 @@ MERGE_SUBSTEPS_CLASSES = { "average": AverageMergeSubstepsSampler, "sample": SampleMergeSubstepsSampler, "sample_uncached": SampleUncachedMergeSubstepsSampler, + "simple": SimpleSubstepsSampler, } diff --git a/py/substep_samplers.py b/py/substep_samplers.py index a5a13dd..6a74c18 100644 --- a/py/substep_samplers.py +++ b/py/substep_samplers.py @@ -11,8 +11,53 @@ from .res_support import _de_second_order from .utils import find_first_unsorted +class SamplerResult: + def __init__( + self, + ss, + sampler, + x, + strength=None, + *, + sigma=None, + sigma_next=None, + s_noise=None, + noise_sampler=None, + final=True, + ): + self.x = x + self.sampler = sampler + self.strength = strength if strength is not None else ss.sigma_up + self.s_noise = s_noise if s_noise is not None else sampler.s_noise + self.sigma = sigma if sigma is not None else ss.sigma + self.sigma_next = sigma_next if sigma_next is not None else ss.sigma_next + self.noise_sampler = noise_sampler if noise_sampler else sampler.noise_sampler + self.final = final + + def get_noise(self, scaled=True): + return self.noise_sampler( + self.sigma, self.sigma_next, out_hw=self.x.shape[-2:] + ).mul_(self.noise_scale if scaled else 1.0) + + @property + def noise_scale(self): + return self.strength * self.s_noise + + def noise_x(self, x=None, scale=1.0): + if x is None: + x = self.x + else: + self.x = x + if self.sigma_next == 0 or self.noise_scale == 0: + return x + self.x = x + self.get_noise() * scale + return self.x + + class SingleStepSampler: name = None + self_noise = 0 + model_calls = 0 def __init__( self, @@ -33,7 +78,7 @@ class SingleStepSampler: self.noise_sampler = noise_sampler self.weight = weight self.substeps = substeps - self.kwargs = kwargs + self.options = kwargs def step(self, x, ss): raise NotImplementedError @@ -43,7 +88,8 @@ class SingleStepSampler: sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) d = to_d(x, ss.sigma, ss.denoised) dt = sigma_down - ss.sigma - return x + d * dt, sigma_up + yield SamplerResult(ss, self, x + d * dt, sigma_up) + # return x + d * dt, sigma_up def __str__(self): return f"" @@ -62,6 +108,9 @@ class SingleStepSampler: def get_dyn_eta(self, ss): return self.eta * self.get_dyn_value(ss, self.dyn_eta_start, self.dyn_eta_end) + def max_noise_samples(self): + return (1 + self.self_noise) * self.substeps + class ReversibleSingleStepSampler(SingleStepSampler): def __init__(self, *, reta=1.0, dyn_reta_start=None, dyn_reta_end=None, **kwargs): @@ -94,17 +143,27 @@ class DPMPPStepBase(SingleStepSampler): class DPMPP2MStep(DPMPPStepBase): def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) t, t_next = self.t_fn(ss.sigma), self.t_fn(ss.sigma_next) h = t_next - t st, st_next = self.sigma_fn(t), self.sigma_fn(t_next) if len(ss.dhist) == 0 or ss.sigma_prev is None: - return (st_next / st) * x - (-h).expm1() * ss.denoised, 0.0 + return ( + yield SamplerResult( + ss, + self, + (st_next / st) * x - (-h).expm1() * ss.denoised, + ss.sigma.new_zeros(1), + ) + ) h_last = t - self.t_fn(ss.sigma_prev) r = h_last / h denoised, old_denoised = ss.denoised, ss.dhist[-1] denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised - return (st_next / st) * x - (-h).expm1() * denoised_d, 0.0 + + yield SamplerResult( + ss, self, (st_next / st) * x - (-h).expm1() * denoised_d, 0.0 + ) class DPMPP2MSDEStep(SingleStepSampler): @@ -116,10 +175,8 @@ class DPMPP2MSDEStep(SingleStepSampler): def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) denoised = ss.denoised - if ss.sigma_next == 0: - return denoised, None # DPM-Solver++(2M) SDE t, s = -ss.sigma.log(), -ss.sigma_next.log() h = s - t @@ -131,7 +188,7 @@ class DPMPP2MSDEStep(SingleStepSampler): ) noise_strength = ss.sigma_next * (-2 * eta_h).expm1().neg().sqrt() if len(ss.dhist) == 0 or ss.sigma_prev is None: - return x, noise_strength + return (yield SamplerResult(ss, self, x, noise_strength)) h_last = (-ss.sigma.log()) - (-ss.sigma_prev.log()) r = h_last / h old_denoised = ss.dhist[-1] @@ -145,7 +202,7 @@ class DPMPP2MSDEStep(SingleStepSampler): x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * ( denoised - old_denoised ) - return x, noise_strength + yield SamplerResult(ss, self, x, noise_strength) class DPMPP3MSDEStep(SingleStepSampler): @@ -153,10 +210,10 @@ class DPMPP3MSDEStep(SingleStepSampler): def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) denoised = ss.denoised - if ss.sigma_next == 0: - return denoised, 0 + # if ss.sigma_next == 0: + # return denoised, 0 t, s = -ss.sigma.log(), -ss.sigma_next.log() h = s - t eta = self.get_dyn_eta(ss) @@ -165,7 +222,7 @@ class DPMPP3MSDEStep(SingleStepSampler): x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised noise_strength = ss.sigma_next * (-2 * h * eta).expm1().neg().sqrt() if len(ss.dhist) == 0 or ss.sigma_prev is None: - return x, noise_strength + return (yield SamplerResult(ss, self, x, noise_strength)) h_1 = (-ss.sigma.log()) - (-ss.sigma_prev.log()) denoised_1 = ss.dhist[-1] if len(ss.dhist) == 1: @@ -185,18 +242,19 @@ class DPMPP3MSDEStep(SingleStepSampler): phi_2 = h_eta.neg().expm1() / h_eta + 1 phi_3 = phi_2 / h_eta - 0.5 x = x + phi_2 * d1 - phi_3 * d2 - return x, noise_strength + yield SamplerResult(ss, self, x, noise_strength) # Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers class ReversibleHeunStep(ReversibleSingleStepSampler): name = "reversible_heun" + model_calls = 1 def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) - sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step( + sigma_down_reversible, _sigma_up_reversible = ss.get_ancestral_step( self.get_dyn_reta(ss) ) dt = sigma_down - ss.sigma @@ -216,19 +274,20 @@ class ReversibleHeunStep(ReversibleSingleStepSampler): # Update the sample using the Reversible Heun formula x = x + dt * (d + d_next) / 2 - dt_reversible**2 * (d_next - d) / 4 - return x, sigma_up + yield SamplerResult(ss, self, x, sigma_up) # Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers class ReversibleHeun1SStep(ReversibleSingleStepSampler): name = "reversible_heun_1s" + model_calls = 1 def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) # Reversible Heun-inspired update (first-order) sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) - sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step( + sigma_down_reversible, _sigma_up_reversible = ss.get_ancestral_step( self.get_dyn_reta(ss) ) sigma_i, sigma_i_plus_1 = ss.sigma, sigma_down @@ -258,12 +317,13 @@ class ReversibleHeun1SStep(ReversibleSingleStepSampler): + dt * (d_i_old + d_i_plus_1) / 2 - dt_reversible**2 * (d_i_plus_1 - d_i_old) / 4 ) - return x, sigma_up + yield SamplerResult(ss, self, x, sigma_up) # Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers class RESStep(SingleStepSampler): name = "res" + model_calls = 1 def __init__(self, *, res_simple_phi=False, res_c2=0.5, **kwargs): super().__init__(**kwargs) @@ -273,7 +333,7 @@ class RESStep(SingleStepSampler): def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) eta = self.get_dyn_eta(ss) sigma_down, sigma_up = ss.get_ancestral_step(eta) denoised = ss.denoised @@ -294,16 +354,17 @@ class RESStep(SingleStepSampler): denoised2 = ss.model(x_2, sigma_2, model_call_idx=1) x = math.exp(-h) * x + h * (b1 * denoised + b2 * denoised2) - return x, sigma_up + yield SamplerResult(ss, self, x, sigma_up) # Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers class TrapezoidalStep(SingleStepSampler): name = "trapezoidal" + model_calls = 1 def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) dt = ss.sigma_next - ss.sigma denoised = ss.denoised @@ -323,19 +384,20 @@ class TrapezoidalStep(SingleStepSampler): dt_2 = sigma_down - ss.sigma # Update the sample using the Trapezoidal rule x = x + dt_2 * (d_i + d_next) / 2 - return x, sigma_up + yield SamplerResult(ss, self, x, sigma_up) # Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers class BogackiStep(ReversibleSingleStepSampler): name = "bogacki" reversible = False + model_calls = 2 def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) - sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step( + sigma_down_reversible, _sigma_up_reversible = ss.get_ancestral_step( self.get_dyn_reta(ss) ) sigma, sigma_next = ss.sigma, sigma_down @@ -370,7 +432,7 @@ class BogackiStep(ReversibleSingleStepSampler): # Update the sample x = x + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9 - correction - return x, sigma_up + yield SamplerResult(ss, self, x, sigma_up) class ReversibleBogackiStep(BogackiStep): @@ -381,10 +443,11 @@ class ReversibleBogackiStep(BogackiStep): # Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers class RK4Step(SingleStepSampler): name = "rk4" + model_calls = 3 def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) sigma = ss.sigma # Calculate the derivative using the model @@ -420,18 +483,19 @@ class RK4Step(SingleStepSampler): # Update the sample x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6 - return x, sigma_up + yield SamplerResult(ss, self, x, sigma_up) # Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers class EulerDancingStep(SingleStepSampler): name = "euler_dancing" + self_noise = 1 def __init__( self, *, deta=1.0, - ds_noise=1.0, + ds_noise=None, leap=2, dyn_deta_start=None, dyn_deta_end=None, @@ -440,7 +504,7 @@ class EulerDancingStep(SingleStepSampler): ): super().__init__(**kwargs) self.deta = deta - self.ds_noise = ds_noise + self.ds_noise = ds_noise if ds_noise is not None else self.s_noise self.leap = leap self.dyn_deta_start = dyn_deta_start self.dyn_deta_end = dyn_deta_end @@ -449,6 +513,51 @@ class EulerDancingStep(SingleStepSampler): self.dyn_deta_mode = dyn_deta_mode def step(self, x, ss): + if ss.sigma_next == 0: + return (yield from self.euler_step(x, ss)) + eta = self.eta + deta = self.deta + leap_sigmas = ss.sigmas[ss.idx :] + leap_sigmas = leap_sigmas[: find_first_unsorted(leap_sigmas)] + zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] + max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 + is_danceable = max_leap > 1 and ss.sigma_next != 0 + curr_leap = max(1, min(self.leap, max_leap)) + sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next + del leap_sigmas + sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) + print("???", sigma_down, sigma_up) + d = to_d(x, ss.sigma, ss.denoised) + # Euler method + dt = sigma_down - ss.sigma + x = x + d * dt + if curr_leap == 1: + return (yield SamplerResult(ss, self, x, sigma_up)) + noise_strength = self.ds_noise * sigma_up + if noise_strength != 0: + x = yield SamplerResult( + ss, self, x, sigma_up, sigma_next=sigma_leap, final=False + ) + # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_( + # self.ds_noise * sigma_up + # ) + # sigma_down2, sigma_up2 = get_ancestral_step(sigma_leap, ss.sigma, eta=deta) + # _sigma_down2, sigma_up2 = get_ancestral_step(sigma_leap, ss.sigma, eta=deta) + # sigma_up2 = ss.sigma_next + (ss.sigma - ss.sigma_next) * 0.5 + sigma_up2 = get_ancestral_step(ss.sigma_next, sigma_leap, eta=deta)[1] + ( + ss.sigma_next * 0.5 + ) + sigma_down2, _sigma_up2 = get_ancestral_step( + ss.sigma_next, sigma_leap, eta=deta + ) + print(">>>", sigma_down2, sigma_up2, "--", ss.sigma, "->", sigma_leap) + # sigma_down2, sigma_up2 = get_ancestral_step(ss.sigma_next, sigma_leap, eta=deta) + d_2 = to_d(x, sigma_leap, ss.denoised) + dt_2 = sigma_down2 - sigma_leap + x = x + d_2 * dt_2 + yield SamplerResult(ss, self, x, sigma_up2) + + def _step(self, x, ss): eta = self.get_dyn_eta(ss) leap_sigmas = ss.sigmas[ss.idx :] leap_sigmas = leap_sigmas[: find_first_unsorted(leap_sigmas)] @@ -457,7 +566,8 @@ class EulerDancingStep(SingleStepSampler): is_danceable = max_leap > 1 and ss.sigma_next != 0 curr_leap = max(1, min(self.leap, max_leap)) sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next - print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas) + # DANCE 35 6 tensor(10.0947, device='cuda:0') -- tensor([21.9220, + # print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas) del leap_sigmas sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) d = to_d(x, ss.sigma, ss.denoised) @@ -467,8 +577,12 @@ class EulerDancingStep(SingleStepSampler): if curr_leap == 1: return x, sigma_up dance_scale = self.get_dyn_value(ss, self.dyn_deta_start, self.dyn_deta_end) - if not is_danceable or abs(dance_scale) < 1e-04: - return x, sigma_up + if curr_leap == 1 or not is_danceable or abs(dance_scale) < 1e-04: + print("NODANCE", dance_scale, self.deta, is_danceable, ss.sigma_next) + yield SamplerResult(ss, self, x, sigma_up) + print( + "DANCE", dance_scale, self.deta, self.dyn_deta_mode, self.ds_noise, sigma_up + ) sigma_down_normal, sigma_up_normal = get_ancestral_step( ss.sigma, ss.sigma_next, eta ) @@ -477,30 +591,133 @@ class EulerDancingStep(SingleStepSampler): x_normal = x + d * dt_normal else: x_normal = x - x = x + self.noise_sampler(ss.sigma, sigma_leap) * self.s_noise * sigma_up sigma_down2, sigma_up2 = get_ancestral_step( sigma_leap, ss.sigma_next, eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), ) + print( + "-->", + sigma_down2, + sigma_up2, + "--", + self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), + ) + x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_(self.ds_noise * sigma_up) d_2 = to_d(x, sigma_leap, ss.denoised) dt_2 = sigma_down2 - sigma_leap result = x + d_2 * dt_2 + # SIGMA: norm_up=9.062416076660156, up=10.703859329223633, up2=19.376544952392578, str=21.955078125 + noise_strength = sigma_up2 + ((sigma_up - sigma_up_normal) ** 5.0) + noise_strength = sigma_up2 + ((sigma_up2 - sigma_up) * 0.5) + # noise_strength = sigma_up2 + ( + # (sigma_up2 - sigma_up) ** (1.0 - (sigma_up_normal / sigma_up2)) + # ) + noise_diff = ( + sigma_up - sigma_up_normal + if sigma_up > sigma_up_normal + else sigma_up_normal - sigma_up + ) + noise_div = ( + sigma_up / sigma_up_normal + if sigma_up > sigma_up_normal + else sigma_up_normal / sigma_up + ) + noise_diff = sigma_up2 - sigma_up_normal + noise_div = sigma_up2 / sigma_up_normal + noise_div = ss.sigma / sigma_leap + + # noise_strength = sigma_up2 + (noise_diff * noise_div) + # noise_strength = sigma_up2 + ((noise_diff * 0.5) ** 2.0) + # noise_strength = sigma_up2 + ((1.0 - noise_diff) ** 0.5) + # noise_strength = sigma_up2 + (((sigma_up2 - sigma_up) * 0.5) ** 2.0) + # noise_strength = sigma_up2 + (((sigma_up2 - sigma_up_normal) * 0.5) ** 1.5) + # noise_strength = sigma_up2 + ( + # (noise_diff * 0.1875) ** (1.0 / (noise_div - 0.0)) + # ) + # noise_strength = sigma_up2 + ( + # (noise_diff * 0.125) ** (1.0 / (noise_div * 1.25)) + # ) + # noise_strength = sigma_up2 + ((noise_diff * 0.2) ** (1.0 / (noise_div * 1.0))) + noise_strength = sigma_up2 + (noise_diff * 0.9 * max(0.0, noise_div - 0.8)) + noise_strength = sigma_up2 + ( + (noise_diff / (curr_leap * 0.4)) + * ((noise_div - (curr_leap / 2.0)).clamp(min=0, max=1.5) * 1.0) + ) + # (1.0 / (noise_div * 1.25))) + # noise_strength = sigma_up2 + ((noise_diff * 0.5) ** noise_div) + print( + f"SIGMA: norm_up={sigma_up_normal}, up={sigma_up}, up2={sigma_up2}, str={noise_strength}", + # noise_diff, + noise_div, + ) + return result, noise_strength + noise_diff = sigma_up2 - sigma_up * dance_scale noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap) + # noise_scale = sigma_up2 * self.ds_noise if self.dyn_deta_mode == "deta" or dance_scale == 1.0: return result, noise_scale result = torch.lerp(x_normal, result, dance_scale) # FIXME: Broken for noise samplers that care about s/sn return result, noise_scale + # def step(self, x, ss): + # eta = self.get_dyn_eta(ss) + # leap_sigmas = ss.sigmas[ss.idx :] + # leap_sigmas = leap_sigmas[: find_first_unsorted(leap_sigmas)] + # zero_idx = (leap_sigmas <= 0).nonzero().flatten()[:1] + # max_leap = (zero_idx.item() if len(zero_idx) else len(leap_sigmas)) - 1 + # is_danceable = max_leap > 1 and ss.sigma_next != 0 + # curr_leap = max(1, min(self.leap, max_leap)) + # sigma_leap = leap_sigmas[curr_leap] if is_danceable else ss.sigma_next + # # print("DANCE", max_leap, curr_leap, sigma_leap, "--", leap_sigmas) + # del leap_sigmas + # sigma_down, sigma_up = get_ancestral_step(ss.sigma, sigma_leap, eta) + # d = to_d(x, ss.sigma, ss.denoised) + # # Euler method + # dt = sigma_down - ss.sigma + # x = x + d * dt + # if curr_leap == 1: + # return x, sigma_up + # dance_scale = self.get_dyn_value(ss, self.dyn_deta_start, self.dyn_deta_end) + # if not is_danceable or abs(dance_scale) < 1e-04: + # print("NODANCE", dance_scale, self.deta) + # return x, sigma_up + # print("NODANCE", dance_scale, self.deta) + # sigma_down_normal, _sigma_up_normal = get_ancestral_step( + # ss.sigma, ss.sigma_next, eta + # ) + # if self.dyn_deta_mode == "lerp": + # dt_normal = sigma_down_normal - ss.sigma + # x_normal = x + d * dt_normal + # else: + # x_normal = x + # x = x + self.noise_sampler(ss.sigma, sigma_leap).mul_(self.s_noise * sigma_up) + # sigma_down2, sigma_up2 = get_ancestral_step( + # sigma_leap, + # ss.sigma_next, + # eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale), + # ) + # d_2 = to_d(x, sigma_leap, ss.denoised) + # dt_2 = sigma_down2 - sigma_leap + # result = x + d_2 * dt_2 + # noise_diff = sigma_up2 - sigma_up * dance_scale + # noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap) + # if self.dyn_deta_mode == "deta" or dance_scale == 1.0: + # return result, noise_scale + # result = torch.lerp(x_normal, result, dance_scale) + # # FIXME: Broken for noise samplers that care about s/sn + # return result, noise_scale + class DPMPP2SStep(DPMPPStepBase): name = "dpmpp_2s" + model_calls = 1 def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) t_fn, sigma_fn = self.t_fn, self.sigma_fn sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) # DPM-Solver++(2S) @@ -511,11 +728,13 @@ class DPMPP2SStep(DPMPPStepBase): x_2 = (sigma_fn(s) / sigma_fn(t)) * x - (-h * r).expm1() * ss.denoised denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=0) x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_2 - return x, sigma_up + yield SamplerResult(ss, self, x, sigma_up) class DPMPPSDEStep(DPMPPStepBase): name = "dpmpp_sde" + self_noise = 1 + model_calls = 1 def __init__(self, *args, r=1 / 2, **kwargs): super().__init__(*args, **kwargs) @@ -523,11 +742,9 @@ class DPMPPSDEStep(DPMPPStepBase): def step(self, x, ss): if ss.sigma_next == 0: - return self.euler_step(x, ss) + return (yield from self.euler_step(x, ss)) t_fn, sigma_fn = self.t_fn, self.sigma_fn - r, eta, s_noise = self.r, self.get_dyn_eta(ss), self.s_noise - noise_sampler = self.noise_sampler - sigma_down, sigma_up = ss.get_ancestral_step(eta) + r, eta = self.r, self.get_dyn_eta(ss) # DPM-Solver++ t, t_next = t_fn(ss.sigma), t_fn(ss.sigma_next) h = t_next - t @@ -538,7 +755,9 @@ class DPMPPSDEStep(DPMPPStepBase): sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta) s_ = t_fn(sd) x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - (t - s_).expm1() * ss.denoised - x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su + x_2 = yield SamplerResult( + ss, self, x_2, su, sigma=sigma_fn(t), sigma_next=sigma_fn(s), final=False + ) denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=1) # Step 2 @@ -546,13 +765,14 @@ class DPMPPSDEStep(DPMPPStepBase): t_next_ = t_fn(sd) denoised_d = (1 - fac) * ss.denoised + fac * denoised_2 x = (sigma_fn(t_next_) / sigma_fn(t)) * x - (t - t_next_).expm1() * denoised_d - return x, su + yield SamplerResult(ss, self, x, su) # Based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers # Which was originally written by Katherine Crowson class TTMJVPStep(SingleStepSampler): name = "ttm_jvp" + model_calls = 1 def __init__(self, *args, alternate_phi_2_calc=True, **kwargs): super().__init__(*args, **kwargs) @@ -560,9 +780,8 @@ class TTMJVPStep(SingleStepSampler): def step(self, x, ss): if ss.sigma_next == 0: - return ss.denoised, ss.sigma.new_zeros(1) + return (yield SamplerResult(ss, self, ss.denoised, ss.sigma.new_zeros(1))) eta = self.get_dyn_eta(ss) - sigma_down, sigma_up = ss.get_ancestral_step(eta) sigma, sigma_next = ss.sigma, ss.sigma_next # 2nd order truncated Taylor method t, s = -sigma.log(), -sigma_next.log() @@ -570,7 +789,7 @@ class TTMJVPStep(SingleStepSampler): h_eta = h * (eta + 1) eps = to_d(x, sigma, ss.denoised) - denoised, denoised_prime = ss.model( + _denoised, denoised_prime = ss.model( x, sigma, tangents=(eps * -sigma, -sigma), model_call_idx=1 ) @@ -582,10 +801,10 @@ class TTMJVPStep(SingleStepSampler): x = torch.exp(-h_eta) * x + phi_1 * ss.denoised + phi_2 * denoised_prime if not eta: - return x, ss.sigma.new_zeros(1) + return (yield SamplerResult(ss, self, x, ss.sigma.new_zeros(1))) phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * eta)) - return x, sigma_next * phi_1_noise + yield SamplerResult(ss, self, x, sigma_next * phi_1_noise) STEP_SAMPLERS = { diff --git a/py/substep_sampling.py b/py/substep_sampling.py index ec2836a..5e307a8 100644 --- a/py/substep_sampling.py +++ b/py/substep_sampling.py @@ -1,15 +1,94 @@ +import gc import torch from comfy.k_diffusion.sampling import get_ancestral_step -class StepSamplerChain: +class Items: def __init__(self, items=None): self.items = [] if items is None else items def clone(self): return self.__class__(items=self.items.copy()) + def append(self, item): + self.items.append(item) + return item + + def __getitem__(self, key): + return self.items[key] + + def __setitem__(self, key, value): + self.items[key] = value + + def __len__(self): + return len(self.items) + + def __iter__(self): + return self.items.__iter__() + + +class CommonOptionsItems(Items): + def __init__(self, *, s_noise=1.0, eta=1.0, items=None, **kwargs): + super().__init__(items=items) + self.options = kwargs + self.s_noise = s_noise + self.eta = eta + + def clone(self): + obj = super().clone() + obj.options = self.options.copy() + obj.s_noise = self.s_noise + obj.eta = self.eta + return obj + + +class StepSamplerChain(CommonOptionsItems): + def __init__( + self, + *, + merge_method="divide", + time_mode="step", + time_start=0, + time_end=999, + **kwargs, + ): + super().__init__(**kwargs) + # step, step_pct, sigma + self.merge_method = merge_method + self.time_mode = time_mode + self.time_start, self.time_end = time_start, time_end + + def check_time(self, sigma, step, steps): + step_pct = step / steps if steps != 0 else 0.0 + if self.time_mode == "step": + return self.time_start <= step <= self.time_end + if self.time_mode == "step_pct": + return self.time_start <= step_pct <= self.time_end + if self.time_mode == "sigma": + return self.time_start >= sigma >= self.time_end + raise ValueError("Bad time mode") + + def clone(self): + obj = super().clone() + obj.merge_method = self.merge_method + obj.time_mode = self.time_mode + obj.time_start, obj.time_end = self.time_start, self.time_end + obj.options = self.options.copy() + return obj + + +class ParamGroup(Items): + pass + + +class StepSamplerGroups(CommonOptionsItems): + def find_match(self, sigma, step, steps): + for idx, item in enumerate(self.items): + if item.check_time(sigma, step, steps): + return idx + return None + class History: def __init__(self, x, size): @@ -37,6 +116,135 @@ class History: self.last = None +class NoiseSamplerCache: + def __init__( + self, + x, + seed, + min_sigma, + max_sigma, + *, + normalize_noise=True, + cpu_noise=True, + batch_size=32, + caching=True, + cache_reset_interval=9999, + set_seed=False, + scale=1.0, + normalize_dims=(-3, -2, -1), + **_unused, + ): + self.x = x + self.mega_x = None + self.seed = seed + self.seed_offset = 0 + self.min_sigma = min_sigma + self.max_sigma = max_sigma + self.cache = {} + self.batch_size = max(1, batch_size) + self.normalize_noise = normalize_noise + self.cpu_noise = cpu_noise + self.caching = caching + self.cache_reset_interval = max(1, cache_reset_interval) + self.scale = float(scale) + self.normalize_dims = tuple(int(v) for v in normalize_dims) + self.update_x(x) + if set_seed: + import random + + random.seed(seed) + torch.manual_seed(seed) + + def reset_cache(self): + self.cache = {} + gc.collect() + + def scale_noise(self, noise, factor=1.0, normalized=None, normalize_dims=None): + normalized = self.normalize_noise if normalized is None else normalized + normalize_dims = ( + self.normalize_dims if normalize_dims is None else normalize_dims + ) + if not normalized or not noise.numel(): + return noise.mul_(factor * self.scale) + mean, std = ( + noise.mean(dim=normalize_dims, keepdim=True), + noise.std(dim=normalize_dims, keepdim=True), + ) + return noise.sub_(mean).div_(std).mul_(factor * self.scale) + + def update_x(self, x): + if self.x.shape == x.shape and self.mega_x is not None: + self.x = x + return + self.x = x + self.cache = {} + self.mega_x = None + if self.batch_size == 1: + self.mega_x = x + return + self.mega_x = x.repeat(x.shape[0] * self.batch_size, *((1,) * (x.dim() - 1))) + + def set_cache(self, key, noise_sampler): + if not self.caching: + return + self.cache[key] = noise_sampler + + def make_caching_noise_sampler(self, nsobj, size, sigma, sigma_next): + size = min(size, self.batch_size) + cache_key = (nsobj, size) + if self.caching: + noise_sampler = self.cache.get(cache_key) + if noise_sampler: + return noise_sampler + curr_seed = self.seed + self.seed_offset + self.seed_offset += 1 + curr_x = self.mega_x[: self.x.shape[0] * size, ...] + if nsobj is None: + + def ns(_s, _sn): + return torch.randn_like(curr_x) + else: + ns = nsobj.make_noise_sampler( + curr_x, + self.min_sigma, + self.max_sigma, + seed=curr_seed, + normalized=False, + cpu=self.cpu_noise, + ) + if self.batch_size == 1: + + def noise_sampler(*_unused, **_unusedkwargs): + return self.scale_noise(ns(sigma, sigma_next)) + + self.set_cache(cache_key, noise_sampler) + return noise_sampler + + orig_h, orig_w = self.x.shape[-2:] + remain = 0 + noise = None + + def noise_sampler(*_unused, out_hw=(orig_h, orig_w)): + nonlocal remain, noise + if out_hw != (orig_h, orig_w): + raise NotImplementedError( + f"Noise size mismatch: {out_hw} vs {(orig_h, orig_w)}" + ) + if remain < 1: + noise = self.scale_noise(ns(sigma, sigma_next)).view( + size, + *self.x.shape, + ) + remain = size + # print("NOISE BATCH", noise.shape, remain) + result = noise[-remain] + remain -= 1 + return result + + self.set_cache(cache_key, noise_sampler) + return noise_sampler + + class ModelCallCache: def __init__( self, model, x, s_in, extra_args, *, size=0, max_use=1000000, threshold=1 @@ -113,6 +321,7 @@ class SamplerState: noise_sampler, callback=None, denoised=None, + noise=None, eta=1.0, reta=1.0, s_noise=1.0, @@ -128,6 +337,7 @@ class SamplerState: self.denoised = denoised self.callback_ = callback self.noise_sampler = noise_sampler + self.noise = noise self.update(idx) def update(self, idx=None): @@ -135,9 +345,9 @@ class SamplerState: self.idx = idx self.sigma_prev = None if idx < 1 else self.sigmas[idx - 1] self.sigma, self.sigma_next = self.sigmas[idx], self.sigmas[idx + 1] - # if self.sigma_prev is not None and self.sigma < self.sigma_prev: - # self.dhist.reset() - # self.xhist.reset() + if self.sigma_prev is not None and self.sigma >= self.sigma_prev: + self.dhist.reset() + self.xhist.reset() self.sigma_down, self.sigma_up = get_ancestral_step( self.sigma, self.sigma_next, eta=self.eta ) @@ -162,6 +372,7 @@ class SamplerState: "denoised", "callback_", "noise_sampler", + "noise", "idx", "sigma", "sigma_next", @@ -178,12 +389,10 @@ class SamplerState: def callback(self, x): if not self.callback_: return None - return self.callback_( - { - "x": x, - "i": self.idx, - "sigma": self.sigma, - "sigma_hat": self.sigma, - "denoised": self.dhist[-1], - } - ) + return self.callback_({ + "x": x, + "i": self.idx, + "sigma": self.sigma, + "sigma_hat": self.sigma, + "denoised": self.dhist[-1], + }) diff --git a/py/utils.py b/py/utils.py index 33f29a8..aeaaf64 100644 --- a/py/utils.py +++ b/py/utils.py @@ -2,18 +2,44 @@ import math import torch -def scale_noise(noise, factor=1.0, *, normalized=True, threshold_std_devs=2.5): +def scale_noise( + noise, + factor=1.0, + *, + normalized=True, + threshold_std_devs=2.5, + normalize_dims=(-3, -2, -1), +): if not normalized or noise.numel() == 0: return noise.mul_(factor) if factor != 1 else noise - mean, std = noise.mean().item(), noise.std().item() - threshold = threshold_std_devs / math.sqrt(noise.numel()) - if abs(mean) > threshold: - noise -= mean - if abs(1.0 - std) > threshold: - noise /= std + mean, std = ( + noise.mean(dim=normalize_dims, keepdim=True), + noise.std(dim=normalize_dims, keepdim=True), + ) + # threshold = threshold_std_devs / math.sqrt(noise.numel()) + # noise[mean.abs() > threshold] -= mean + # noise[(1.0 - std).abs() > threshold] /= std + noise -= mean + noise /= std + # if abs(mean) > threshold: + # noise -= mean + # if abs(1.0 - std) > threshold: + # noise /= std return noise.mul_(factor) if factor != 1 else noise +# def scale_noise(noise, factor=1.0, *, normalized=True, threshold_std_devs=2.5): +# if not normalized or noise.numel() == 0: +# return noise.mul_(factor) if factor != 1 else noise +# mean, std = noise.mean().item(), noise.std().item() +# threshold = threshold_std_devs / math.sqrt(noise.numel()) +# if abs(mean) > threshold: +# noise -= mean +# if abs(1.0 - std) > threshold: +# noise /= std +# return noise.mul_(factor) if factor != 1 else noise + + def find_first_unsorted(tensor, desc=True): if not (len(tensor.shape) and tensor.shape[0]): return None