From aa3e12ae342c73c0c6a57bfe3dc7ca6020b98bf8 Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 30 May 2024 07:47:43 -0600 Subject: [PATCH] Initial implementation --- README.md | 76 ++++++- __init__.py | 8 + py/__init__.py | 0 py/nodes.py | 144 +++++++++++++ py/res_support.py | 127 +++++++++++ py/sampling.py | 88 ++++++++ py/substep_merging.py | 252 ++++++++++++++++++++++ py/substep_samplers.py | 469 +++++++++++++++++++++++++++++++++++++++++ py/substep_sampling.py | 144 +++++++++++++ py/utils.py | 22 ++ 10 files changed, 1328 insertions(+), 2 deletions(-) create mode 100644 __init__.py create mode 100644 py/__init__.py create mode 100644 py/nodes.py create mode 100644 py/res_support.py create mode 100644 py/sampling.py create mode 100644 py/substep_merging.py create mode 100644 py/substep_samplers.py create mode 100644 py/substep_sampling.py create mode 100644 py/utils.py diff --git a/README.md b/README.md index cdeb12f..7927e92 100644 --- a/README.md +++ b/README.md @@ -1,2 +1,74 @@ -# comfyui_overly_complicated_sampling -Wildly unsound and experimental sampling for ComfyUI +# Overly Complicated Sampling +Wildly unsound and experimental sampling for [ComfyUI](https://github.com/comfyanonymous/ComfyUI). + +## Description + +Very unstable, experimental and mathematically unsound sampling for ComfyUI. + +Current status: In flux, not suitable for general use. + +## Nodes + +### ComposableSampler + +**Possible Parameters** + +* `avgmerge_stretch`(`0.4`): Used for `average` and `sample` merge types. See below. +* `model_call_cache`(unset): Caches the result of model calls at n+1 (where `n` is the number of model evaluations per step). For example, Bogacki is three model calls per step: whether the first one runs is dependent on the merge strategy. After that, Bogacki calls the model two more times. If you set `model_call_cache` to `1` then the result of that second call will be cached and if you're running two Bogacki substeps then the second one will use the cached version. Massively accelerates inference (especially when using the `average` merge strategy) but is likely very unsound and inaccurate. + +#### Merging + +When running multiple substeps per step, the results will combined based on the merge strategy. Possible strategies: + +* `normal`: The model is called at least once per substep (and possibly additional times for higher order samplers). The result of each substep is noised and the next substep uses that result. Then all the results are averaged. +* `divide`: Creates a linear schedule between the current sigma and the next and runs the substeps in sequence. The model is called at least once per substep. +* `average`: The model is called once at the beginning of the step and substeps share that result (but it may be called additional times for higher order samplers). This means substeps for samplers like reversible Euler, Heun 1s, DPM++ 2m SDE are essentially free. May be theoretically very unsound and inaccurate, requires manual tweaking of settings like `s_noise`. Supports the parameter `avgmerge_stretch`(`0.4`) which basically rolls back the current sigma and adds some noise (otherwise running a substep is deterministic and there would be no point to running a sampler like Euler more than once). +* `sample`: Like `average` (and uses `avgmerge_stretch`) but instead of simply using the average, it does a sampler step toward it instead. You can plug in any substep sampler to the `merge_sampler_opt` input (if unconnected and the merge method is `sample` then Euler will be used). + +### ComposableStepSampler + +This node has a text input for YAML (or JSON) advanced parameters. + +For example, you could enter something like this in the field: + +```yaml +reta: 1.1 +leap: 3 +dyn_deta_mode: "deta" +``` + +**Possible Parameters** + +#### General + +* `eta`(`1.0`): Will override `eta` in the node if set. +* `s_noise`(`1.0`): Will override `s_noise` in the node if set. +* `solver_type`(`midpoint`): Applies to DPM++ 2m SDE. May be one of `midpoint` or `heun` (`midpoint` is generally recommended). + +#### Reversible + +* `reta`(`1.0`): Reverse ETA. + +#### Dancing + +* `leap`(`2`): Distance to try to leap forward. If you set `leap` to `1` you just get plain old Euler ancestral. +* `deta`(`1.0`): ETA used for dance steps. +* `dyn_deta_start`(`unset`) and `dyn_deta_end`(`unset`): No effect unless both values are set. Will interpolate between start and end based on the percentage of sampling. +* `dyn_deta_mode`(`lerp`): May be one of: + * `deta`: Scales `deta` based on the value from `dyn_deta_start/end`. + * `lerp`: Does the dance step according to `deta` and then LERPs the non-dance sample result with the dance sample result based on the scale calculated from `dyn_deta_start/end` (which is `1.0` if they are unset). For example, if the dance scale is `0.5` you will get 50% normal sampling, 50% dancing sampling. + +#### RES + +* `res_simple_phi`(`false`): Applies to RES. Uses a faster but possibly less accurate method for calculating phi. What does phi do? I haven't the foggiest! +* `res_c2`(`0.5`): Applies to RES. Solver partial step size, the default of `0.5` appears to use the midpoint. Setting it to a lower value might possibly be more accurate but slower? + +## Credits + +I can move code around but sampling math and creating samplers is far beyond my ability. I didn't write any of the original samplers: + +* Euler, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation. +* Reversible Heun, Reversible Heun 1s, RES, Trapezoidal, Bogacki, Reversible Bogacki, RK4 and Euler Dancing samplers based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +* Normal substep merge strategy based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers + +Thanks! diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..f7e3aba --- /dev/null +++ b/__init__.py @@ -0,0 +1,8 @@ +from .py import nodes + + +NODE_CLASS_MAPPINGS = { + "ComposableSampler": nodes.ComposableSampler, + "ComposableStepSampler": nodes.ComposableStepSampler, +} +__all__ = ["NODE_CLASS_MAPPINGS"] diff --git a/py/__init__.py b/py/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/py/nodes.py b/py/nodes.py new file mode 100644 index 0000000..dbbea5a --- /dev/null +++ b/py/nodes.py @@ -0,0 +1,144 @@ +from .sampling import composable_sampler, STEP_SAMPLERS +from .substep_sampling import StepSamplerChain +from .substep_merging import MERGE_SUBSTEPS_CLASSES + +import comfy +import yaml + + +class ComposableSampler: + RETURN_TYPES = ("SAMPLER",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "go" + + @classmethod + 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",), + }, + "optional": { + "merge_sampler_opt": ("STEP_SAMPLER_CHAIN",), + "parameters": ( + "STRING", + {"default": "", "multiline": True, "dynamicPrompts": False}, + ), + }, + } + + def go( + self, + *, + s_noise, + eta, + merge_method, + step_sampler_chain, + merge_sampler_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, + } + 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() + return ( + comfy.samplers.KSAMPLER( + composable_sampler, + {"composable_sampler_options": options}, + ), + ) + + +class ComposableStepSampler: + RETURN_TYPES = ("STEP_SAMPLER_CHAIN",) + CATEGORY = "sampling/custom_sampling/samplers" + + FUNCTION = "go" + + @classmethod + 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": 100}), + "step_method": (tuple(STEP_SAMPLERS.keys()),), + }, + "optional": { + "step_sampler_opt": ("STEP_SAMPLER_CHAIN",), + "custom_noise_opt": ("SONAR_CUSTOM_NOISE",), + "parameters": ( + "STRING", + {"default": "", "multiline": True, "dynamicPrompts": False}, + ), + }, + } + + def go(self, *, parameters="", step_sampler_opt=None, **kwargs): + if step_sampler_opt is not None: + chain = step_sampler_opt.clone() + else: + chain = StepSamplerChain() + 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) + return (chain,) + + +__all__ = ("ComposableStepSampler", "ComposableSampler") diff --git a/py/res_support.py b/py/res_support.py new file mode 100644 index 0000000..788f1ec --- /dev/null +++ b/py/res_support.py @@ -0,0 +1,127 @@ +import math + +import torch + +from torch import FloatTensor +from typing import Optional, NamedTuple + +# Copied from https://github.com/Clybius/ComfyUI-Extra-Samplers + + +def _gamma( + n: int, +) -> int: + """ + https://en.wikipedia.org/wiki/Gamma_function + for every positive integer n, + Γ(n) = (n-1)! + """ + return math.factorial(n - 1) + + +def _incomplete_gamma(s: int, x: float, gamma_s: Optional[int] = None) -> float: + """ + https://en.wikipedia.org/wiki/Incomplete_gamma_function#Special_values + if s is a positive integer, + Γ(s, x) = (s-1)!*∑{k=0..s-1}(x^k/k!) + """ + if gamma_s is None: + gamma_s = _gamma(s) + + sum_: float = 0 + # {k=0..s-1} inclusive + for k in range(s): + numerator: float = x**k + denom: int = math.factorial(k) + quotient: float = numerator / denom + sum_ += quotient + incomplete_gamma_: float = sum_ * math.exp(-x) * gamma_s + return incomplete_gamma_ + + +# by Katherine Crowson +def _phi_1(neg_h: FloatTensor): + return torch.nan_to_num(torch.expm1(neg_h) / neg_h, nan=1.0) + + +# by Katherine Crowson +def _phi_2(neg_h: FloatTensor): + return torch.nan_to_num((torch.expm1(neg_h) - neg_h) / neg_h**2, nan=0.5) + + +# by Katherine Crowson +def _phi_3(neg_h: FloatTensor): + return torch.nan_to_num( + (torch.expm1(neg_h) - neg_h - neg_h**2 / 2) / neg_h**3, nan=1 / 6 + ) + + +def _phi( + neg_h: float, + j: int, +): + """ + For j={1,2,3}: you could alternatively use Kat's phi_1, phi_2, phi_3 which perform fewer steps + + Lemma 1 + https://arxiv.org/abs/2308.02157 + ϕj(-h) = 1/h^j*∫{0..h}(e^(τ-h)*(τ^(j-1))/((j-1)!)dτ) + + https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84 + = 1/h^j*[(e^(-h)*(-τ)^(-j)*τ(j))/((j-1)!)]{0..h} + https://www.wolframalpha.com/input?i=integrate+e%5E%28%CF%84-h%29*%28%CF%84%5E%28j-1%29%2F%28j-1%29%21%29d%CF%84+between+0+and+h + = 1/h^j*((e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h)))/(j-1)!) + = (e^(-h)*(-h)^(-j)*h^j*(Γ(j)-Γ(j,-h))/((j-1)!*h^j) + = (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/(j-1)! + = (e^(-h)*(-h)^(-j)*(Γ(j)-Γ(j,-h))/Γ(j) + = (e^(-h)*(-h)^(-j)*(1-Γ(j,-h)/Γ(j)) + + requires j>0 + """ + assert j > 0 + gamma_: float = _gamma(j) + incomp_gamma_: float = _incomplete_gamma(j, neg_h, gamma_s=gamma_) + + phi_: float = math.exp(neg_h) * neg_h**-j * (1 - incomp_gamma_ / gamma_) + + return phi_ + + +class RESDECoeffsSecondOrder(NamedTuple): + a2_1: float + b1: float + b2: float + + +def _de_second_order( + h: float, + c2: float, + simple_phi_calc=False, +) -> RESDECoeffsSecondOrder: + """ + Table 3 + https://arxiv.org/abs/2308.02157 + ϕi,j := ϕi,j(-h) = ϕi(-cj*h) + a2_1 = c2ϕ1,2 + = c2ϕ1(-c2*h) + b1 = ϕ1 - ϕ2/c2 + """ + if simple_phi_calc: + # Kat computed simpler expressions for phi for cases j={1,2,3} + a2_1: float = c2 * _phi_1(-c2 * h) + phi1: float = _phi_1(-h) + phi2: float = _phi_2(-h) + else: + # I computed general solution instead. + # they're close, but there are slight differences. not sure which would be more prone to numerical error. + a2_1: float = c2 * _phi(j=1, neg_h=-c2 * h) + phi1: float = _phi(j=1, neg_h=-h) + phi2: float = _phi(j=2, neg_h=-h) + phi2_c2: float = phi2 / c2 + b1: float = phi1 - phi2_c2 + b2: float = phi2_c2 + return RESDECoeffsSecondOrder( + a2_1=a2_1, + b1=b1, + b2=b2, + ) diff --git a/py/sampling.py b/py/sampling.py new file mode 100644 index 0000000..da58967 --- /dev/null +++ b/py/sampling.py @@ -0,0 +1,88 @@ +import torch +from tqdm.auto import trange + + +from .substep_samplers import STEP_SAMPLERS +from .substep_sampling import SamplerState, History +from .substep_merging import MERGE_SUBSTEPS_CLASSES + + +def composable_sampler( + model, + x, + sigmas, + *, + s_noise=1.0, + eta=1.0, + composable_sampler_options, + extra_args=None, + callback=None, + disable=None, + noise_sampler=None, + **kwargs, +): + copts = composable_sampler_options.copy() + if extra_args is None: + extra_args = {} + if noise_sampler is None: + + def noise_sampler(_s, _sn): + return torch.randn_like(x) + + ss = SamplerState( + model, + sigmas, + 0, + x.new_ones((x.shape[0],)), + History(x, 3), + History(x, 2), + extra_args, + model_call_cache=None + if "model_call_cache" not in copts + else History(x, copts["model_call_cache"]), + noise_sampler=noise_sampler, + callback=callback, + eta=eta if eta != 1.0 else copts["eta"], + s_noise=s_noise if s_noise != 1.0 else copts["s_noise"], + reta=copts.get("reta", 1.0), + ) + 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 += (ssampler,) * sitem["substeps"] + substeps += sitem["substeps"] + msitem = copts["merge_sampler"] + if copts["merge_method"] == "sample": + 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 + merge_sampler = MERGE_SUBSTEPS_CLASSES[copts["merge_method"]]( + ss, + samplers, + **(copts | {"merge_sampler": merge_sampler}), + ) + for idx in trange(len(sigmas) - 1, disable=disable): + print(f"STEP {idx+1}") + ss.update(idx) + if ss.model_call_cache is not None: + ss.model_call_cache.reset() + x = merge_sampler.step(x) + return x diff --git a/py/substep_merging.py b/py/substep_merging.py new file mode 100644 index 0000000..c26aefa --- /dev/null +++ b/py/substep_merging.py @@ -0,0 +1,252 @@ +import torch + +from .utils import scale_noise, find_first_unsorted +from .substep_sampling import History + + +class MergeSubstepsSampler: + def __init__(self, ss, samplers, **_kwargs): + self.ss = ss + self.samplers = samplers + + def step(self, x): + raise NotImplementedError + + def merge_steps(self, _x, result): + return result + + +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 = len(self.samplers) + renoise_weight = 1.0 / substeps + z_avg = torch.zeros_like(x) + noise = torch.zeros_like(x) + noise_total = 0.0 + for subidx, ssampler in enumerate(self.samplers): + print(" SUBSTEP", subidx, ssampler) + ss.denoised = ss.model(x, ss.sigma * ss.s_in) + z_k, noise_strength = ( + ssampler.step if ss.sigma_next != 0 else ssampler.final_step + )(x, ss) + z_avg += renoise_weight * z_k + if ss.sigma_next == 0: + continue + noise_strength *= ssampler.s_noise + noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next) + x = z_k + if subidx != 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 + + +class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler): + def __init__(self, ss, samplers, *, avgmerge_stretch=0.4, **kwargs): + super().__init__(ss, samplers, **kwargs) + self.ss = ss + self.stretch = avgmerge_stretch + + def step(self, x): + ss = orig_ss = self.ss + substeps = len(self.samplers) + renoise_weight = 1.0 / substeps + z_avg = torch.zeros_like(x) + noise = torch.zeros_like(x) + stretch = (ss.sigma - ss.sigma_next) * self.stretch + sig_adj = ss.sigma + stretch + ss = self.ss.clone_edit(sigma=sig_adj) + x = x + ss.noise_sampler(orig_ss.sigma, ss.sigma_next) * stretch * ss.s_noise + ss.denoised = ss.model(x, sig_adj * ss.s_in) + noise_total = 0.0 + for subidx, ssampler in enumerate(self.samplers): + print(" SUBSTEP", subidx, ssampler) + curr_x = x + scale_noise( + ssampler.noise_sampler(sig_adj, ss.sigma_next), + ssampler.s_noise * stretch, + ) + z_k, noise_strength = ( + ssampler.step if ss.sigma_next != 0 else ssampler.final_step + )(curr_x, ss) + z_avg += renoise_weight * z_k + if ss.sigma_next == 0: + continue + noise_strength *= ssampler.s_noise + noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next) + 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 + + +class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): + def __init__( + self, ss, samplers, *, merge_sample_skip=False, merge_sampler, **kwargs + ): + super().__init__(ss, samplers, **kwargs) + self.merge_sampler = merge_sampler + self.merge_ss = None + self.merge_sample_skip = merge_sample_skip + + def step(self, x): + ss = self.ss + substeps = len(self.samplers) + renoise_weight = 1.0 / substeps + z_avg = torch.zeros_like(x) + curr_x = x + ss.denoised = None + stretch = (ss.sigma - ss.sigma_next) * self.stretch + sig_adj = ss.sigma + stretch + if self.merge_sample_skip: + stretch = (ss.sigma - ss.sigma_next) * self.stretch + sig_adj = ss.sigma + stretch + ss.denoised = ss.model( + curr_x, + # + ss.noise_sampler(sig_adj.sigma, ss.sigma_next) * stretch * ss.s_noise, + ss.s_in * sig_adj, + ) + for subidx, ssampler in enumerate(self.samplers): + if not self.merge_sample_skip: + ss.denoised = ss.model(curr_x, ss.s_in * ss.sigma) + else: + curr_x = ( + x + + ssampler.noise_sampler(sig_adj, ss.sigma_next) + * ssampler.s_noise + * stretch + ) + print(" SUBSTEP", subidx, ssampler) + z_k, noise_strength = ( + ssampler.step if ss.sigma_next != 0 else ssampler.final_step + )(curr_x, ss) + z_avg += renoise_weight * z_k + curr_x = z_k + if not noise_strength or ss.sigma_next == 0 or self.merge_sample_skip: + continue + curr_x += ( + ssampler.noise_sampler(ss.sigma, ss.sigma_next) + * ssampler.s_noise + * noise_strength + ) + 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 + + def merge_steps(self, x, result): + if self.ss.model_call_cache is not None: + self.ss.model_call_cache.reset() + msampler = self.merge_sampler + if self.merge_ss is None: + merge_ss = self.merge_ss = self.ss.clone_edit( + denoised=result, + dhist=History(x, 3), + xhist=History(x, 2), + s_noise=msampler.s_noise, + eta=msampler.eta, + ) + 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 if not final else msampler.final_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 + ) + merge_ss.dhist.push(result) + merge_ss.xhist.push(merged) + merge_ss.denoised = None + return merged + + +class DivideMergeSubstepsSampler(MergeSubstepsSampler): + def __init__(self, ss, samplers, *, schedule_multiplier=4, **kwargs): + super().__init__(ss, samplers, **kwargs) + self.schedule_multiplier = schedule_multiplier + + def step(self, x): + ss = self.ss + samplers = self.samplers + substeps = len(samplers) + max_steps = len(self.ss.sigmas) - 1 + sigmas_slice = ss.sigmas[ + ss.idx : min(max_steps + 1, ss.idx + self.schedule_multiplier) + ] + print("SLICE", sigmas_slice) + unsorted_idx = find_first_unsorted(sigmas_slice) + if unsorted_idx is not None: + sigmas_slice = sigmas_slice[:unsorted_idx] + print("SLICE ADJ", sigmas_slice) + chunks = tuple( + torch.linspace( + sigmas_slice[idx], + sigmas_slice[idx + 1], + steps=substeps + 1, + device=sigmas_slice.device, + dtype=sigmas_slice.dtype, + )[0 if not idx else 1 :] + for idx in range(len(sigmas_slice) - 1) + ) + print("CHUNKS", chunks) + subsigmas = torch.cat(chunks) + print("SUBSIGMAS", subsigmas) + subss = self.ss.clone_edit(idx=0, sigmas=subsigmas) + subss.main_idx = ss.idx + subss.main_sigmas = ss.sigmas + for subidx in range(substeps): + subss.update(subidx) + subss.denoised = subss.model(x, subss.sigma) + ssampler = samplers[subidx] + x, noise_strength = ( + ssampler.step if subss.sigma_next != 0 else ssampler.final_step + )(x, subss) + if not noise_strength or subss.sigma_next == 0: + continue + x = ( + x + + ssampler.noise_sampler(subss.sigma, subss.sigma_next) + * ssampler.s_noise + * noise_strength + ) + subss.xhist.push(x) + subss.dhist.push(subss.denoised) + subss.denoised = None + ss.callback(x) + return x + + +MERGE_SUBSTEPS_CLASSES = { + "normal": NormalMergeSubstepsSampler, + "average": AverageMergeSubstepsSampler, + "sample": SampleMergeSubstepsSampler, + "divide": DivideMergeSubstepsSampler, +} diff --git a/py/substep_samplers.py b/py/substep_samplers.py new file mode 100644 index 0000000..491e9ee --- /dev/null +++ b/py/substep_samplers.py @@ -0,0 +1,469 @@ +import math + +import torch + +from comfy.k_diffusion.sampling import ( + get_ancestral_step, + to_d, +) + +from .res_support import _de_second_order +from .utils import find_first_unsorted + + +class SingleStepSampler: + name = None + + def __init__( + self, *, noise_sampler=None, s_noise=1.0, eta=1.0, weight=1.0, **kwargs + ): + self.s_noise = s_noise + self.eta = eta + self.noise_sampler = noise_sampler + self.weight = weight + self.kwargs = kwargs + + def step(self, x, ss): + raise NotImplementedError + + # Euler - based on original ComfyUI implementation + def final_step(self, x, ss): + sigma_down, sigma_up = ss.get_ancestral_step(self.eta) + d = to_d(x, ss.sigma, ss.denoised) + dt = sigma_down - ss.sigma + return x + d * dt, sigma_up + + def __str__(self): + return f"" + + +class ReversibleSingleStepSampler(SingleStepSampler): + def __init__(self, *, reta=1.0, **kwargs): + super().__init__(**kwargs) + self.reta = reta + + +class EulerStep(SingleStepSampler): + name = "euler" + + def step(self, x, ss): + return self.final_step(x, ss) + + +class DPMPP2MStep(SingleStepSampler): + @staticmethod + def sigma_fn(t): + return t.neg().exp() + + @staticmethod + def t_fn(t): + return t.log().neg() + + def step(self, 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 + 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 + + +class DPMPP2MSDEStep(SingleStepSampler): + name = "dpmpp_2m_sde" + + def __init__(self, *, solver_type="midpoint", **kwargs): + super().__init__(**kwargs) + self.solver_type = solver_type + + def step(self, 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 + eta_h = self.eta * h + + x = ( + ss.sigma_next / ss.sigma * (-eta_h).exp() * x + + (-h - eta_h).expm1().neg() * denoised + ) + 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 + h_last = (-ss.sigma.log()) - (-ss.sigma_prev.log()) + r = h_last / h + old_denoised = ss.dhist[-1] + if self.solver_type == "heun": + x = x + ( + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) + * (1 / r) + * (denoised - old_denoised) + ) + elif self.solver_type == "midpoint": + x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * ( + denoised - old_denoised + ) + return x, noise_strength + + +class DPMPP3MSDEStep(SingleStepSampler): + name = "dpmpp_3m_sde" + + def step(self, x, ss): + denoised = ss.denoised + if ss.sigma_next == 0: + return denoised, 0 + t, s = -ss.sigma.log(), -ss.sigma_next.log() + h = s - t + h_eta = h * (self.eta + 1) + + x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised + noise_strength = ss.sigma_next * (-2 * h * self.eta).expm1().neg().sqrt() + if len(ss.dhist) == 0 or ss.sigma_prev is None: + return x, noise_strength + h_1 = (-ss.sigma.log()) - (-ss.sigma_prev.log()) + denoised_1 = ss.dhist[-1] + if len(ss.dhist) == 1: + r = h_1 / h + d = (denoised - denoised_1) / r + phi_2 = h_eta.neg().expm1() / h_eta + 1 + x = x + phi_2 * d + else: + h_2 = (-ss.sigma_prev.log()) - (-ss.sigmas[ss.idx - 2].log()) + denoised_2 = ss.dhist[-2] + r0 = h_1 / h + r1 = h_2 / h + d1_0 = (denoised - denoised_1) / r0 + d1_1 = (denoised_1 - denoised_2) / r1 + d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1) + d2 = (d1_0 - d1_1) / (r0 + r1) + 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 + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class ReversibleHeunStep(ReversibleSingleStepSampler): + name = "reversible_heun" + + def step(self, x, ss): + sigma_down, sigma_up = ss.get_ancestral_step(self.eta) + sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(self.reta) + dt = sigma_down - ss.sigma + dt_reversible = sigma_down_reversible - ss.sigma + + # Calculate the derivative using the model + d = to_d(x, ss.sigma, ss.denoised) + + # Predict the sample at the next sigma using Euler step + x_pred = x + d * dt + + # Denoised sample at the next sigma + denoised_next = ss.model(x_pred, sigma_down, model_call_idx=0) + + # Calculate the derivative at the next sigma + d_next = to_d(x_pred, sigma_down, denoised_next) + + # 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 + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class ReversibleHeun1SStep(ReversibleSingleStepSampler): + name = "reversible_heun_1s" + + def step(self, x, ss): + # Reversible Heun-inspired update (first-order) + sigma_down, sigma_up = ss.get_ancestral_step(self.eta) + sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(self.reta) + sigma_i, sigma_i_plus_1 = ss.sigma, sigma_down + dt = sigma_i_plus_1 - sigma_i + dt_reversible = sigma_down_reversible - sigma_i + + eff_x = ss.xhist[-1] if len(ss.xhist) else x + + # Calculate the derivative using the model + print("Can skip", len(ss.dhist)) + d_i_old = to_d( + eff_x, + sigma_i, + ss.dhist[-1] + if len(ss.dhist) + else ss.model(eff_x, sigma_i, model_call_idx=0), + ) + + # Predict the sample at the next sigma using Euler step + x_pred = eff_x + d_i_old * dt + + # Calculate the derivative at the next sigma + d_i_plus_1 = to_d(x_pred, sigma_i_plus_1, ss.denoised) + + # Update the sample using the Reversible Heun formula + x = ( + x + + dt * (d_i_old + d_i_plus_1) / 2 + - dt_reversible**2 * (d_i_plus_1 - d_i_old) / 4 + ) + return x, sigma_up + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class RESStep(SingleStepSampler): + name = "res" + + def __init__(self, *, res_simple_phi=False, res_c2=0.5, **kwargs): + super().__init__(**kwargs) + self.simple_phi = res_simple_phi + self.c2 = res_c2 + pass + + def step(self, x, ss): + sigma_down, sigma_up = ss.get_ancestral_step(self.eta) + denoised = ss.denoised + lam_next = ( + sigma_down.log().neg() if self.eta != 0 else ss.sigma_next.log().neg() + ) + lam = ss.sigma.log().neg() + + h = lam_next - lam + a2_1, b1, b2 = _de_second_order( + h=h, c2=self.c2, simple_phi_calc=self.simple_phi + ) + + c2_h = 0.5 * h + + x_2 = math.exp(-c2_h) * x + a2_1 * h * denoised + lam_2 = lam + c2_h + sigma_2 = lam_2.neg().exp() + + denoised2 = ss.model(x_2, sigma_2, model_call_idx=0) + + x = math.exp(-h) * x + h * (b1 * denoised + b2 * denoised2) + return x, sigma_up + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class TrapezoidalStep(SingleStepSampler): + name = "trapezoidal" + + def step(self, x, ss): + sigma_down, sigma_up = ss.get_ancestral_step(self.eta) + dt = ss.sigma_next - ss.sigma + denoised = ss.denoised + + # Calculate the derivative using the model + d_i = to_d(x, ss.sigma, denoised) + + # Predict the sample at the next sigma using Euler step + x_pred = x + d_i * dt + + # Denoised sample at the next sigma + denoised_next = ss.model(x_pred, ss.sigma_next, model_call_idx=0) + + # Calculate the derivative at the next sigma + d_next = to_d(x_pred, ss.sigma_next, denoised_next) + + 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 + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class BogackiStep(ReversibleSingleStepSampler): + name = "bogacki" + reversible = False + + def step(self, x, ss): + sigma_down, sigma_up = ss.get_ancestral_step(self.eta) + sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(self.reta) + sigma, sigma_next = ss.sigma, sigma_down + dt = sigma_next - sigma + dt_reversible = sigma_down_reversible - sigma + denoised = ss.denoised + + # Calculate the derivative using the model + d = to_d(x, sigma, denoised) + + # Bogacki-Shampine steps + k1 = d * dt + k2 = ( + to_d( + x + k1 / 2, + sigma + dt / 2, + ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=0), + ) + * dt + ) + k3 = ( + to_d( + x + 3 * k1 / 4 + k2 / 4, + sigma + 3 * dt / 4, + ss.model(x + 3 * k1 / 4 + k2 / 4, sigma + 3 * dt / 4, model_call_idx=1), + ) + * dt + ) + + # Reversible correction term (inspired by Reversible Heun) + correction = dt_reversible**2 * (k3 - k2) / 6 if self.reversible else 0.0 + + # Update the sample + x = x + 2 * k1 / 9 + k2 / 3 + 4 * k3 / 9 - correction + return x, sigma_up + + +class ReversibleBogackiStep(BogackiStep): + name = "reversible_bogacki" + reversible = True + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class RK4Step(SingleStepSampler): + name = "rk4" + + def step(self, x, ss): + sigma_down, sigma_up = ss.get_ancestral_step(self.eta) + sigma = ss.sigma + # Calculate the derivative using the model + d = to_d(x, sigma, ss.denoised) + dt = sigma_down - sigma + + # Runge-Kutta steps + k1 = d * dt + k2 = ( + to_d( + x + k1 / 2, + sigma + dt / 2, + ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=0), + ) + * dt + ) + k3 = ( + to_d( + x + k2 / 2, + sigma + dt / 2, + ss.model(x + k2 / 2, sigma + dt / 2, model_call_idx=1), + ) + * dt + ) + k4 = ( + to_d( + x + k3, + sigma + dt, + ss.model(x + k3, sigma + dt, model_call_idx=2), + ) + * dt + ) + + # Update the sample + x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6 + return x, sigma_up + + +# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers +class EulerDancingStep(SingleStepSampler): + name = "euler_dancing" + + def __init__( + self, + *, + deta=1.0, + ds_noise=1.0, + leap=2, + dyn_deta_start=None, + dyn_deta_end=None, + dyn_deta_mode="lerp", + **kwargs, + ): + super().__init__(**kwargs) + self.deta = deta + self.ds_noise = ds_noise + self.leap = leap + self.dyn_deta_start = dyn_deta_start + self.dyn_deta_end = dyn_deta_end + if dyn_deta_mode not in ("lerp", "deta"): + raise ValueError("Bad dyn_deta_mode") + self.dyn_deta_mode = dyn_deta_mode + + def step(self, x, 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, self.eta) + d = to_d(x, ss.sigma, ss.denoised) + # Euler method + dt = sigma_down - ss.sigma + x = x + d * dt + if None not in (self.dyn_deta_start, self.dyn_deta_end): + if self.dyn_deta_start == self.dyn_deta_end: + dance_scale = self.dyn_deta_start + else: + main_idx = getattr(ss, "main_idx", ss.idx) + main_sigmas = getattr(ss, "main_sigmas", ss.sigmas) + step_pct = main_idx / (len(main_sigmas) - 1) + dd_diff = self.dyn_deta_end - self.dyn_deta_start + dance_scale = self.dyn_deta_start + dd_diff * step_pct + else: + dance_scale = 1.0 + print("DANCE?", dance_scale, ss.idx, is_danceable, curr_leap) + if not is_danceable or abs(dance_scale) < 1e-04: + return x, sigma_up + orig_x = 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 == "lerp" else dance_scale), + ) + d_2 = to_d(x, sigma_leap, ss.denoised) + dt_2 = sigma_down2 - sigma_leap + result = x + d_2 * dt_2 + if self.dyn_deta_mode == "deta" or dance_scale == 1.0: + return result, sigma_up2 + result = torch.lerp(orig_x, result, dance_scale) + # FIXME: Broken for noise samplers that care about s/sn + return result, torch.lerp(sigma_up, sigma_up2, dance_scale) + + +STEP_SAMPLERS = { + "euler": EulerStep, + "dpmpp_2m": DPMPP2MStep, + "dpmpp_2m_sde": DPMPP2MSDEStep, + "dpmpp_3m_sde": DPMPP3MSDEStep, + "reversible_heun": ReversibleHeunStep, + "reversible_heun_1s": ReversibleHeun1SStep, + "res": RESStep, + "trapezoidal": TrapezoidalStep, + "bogacki": BogackiStep, + "reversible_bogacki": ReversibleBogackiStep, + "rk4": RK4Step, + "euler_dancing": EulerDancingStep, +} + +__all__ = ( + "STEP_SAMPLERS", + "EulerStep", + "DPMPP2MStep", + "DPMPP2MSDEStep", + "DPMPP3MSDEStep", + "ReversibleHeunStep", + "ReversibleHeun1SStep", + "RESStep", + "TrapezoidalStep", + "BogackiStep", + "ReversibleBogackiStep", + "EulerDancingStep", +) diff --git a/py/substep_sampling.py b/py/substep_sampling.py new file mode 100644 index 0000000..aa056df --- /dev/null +++ b/py/substep_sampling.py @@ -0,0 +1,144 @@ +import torch + +from comfy.k_diffusion.sampling import get_ancestral_step + + +class StepSamplerChain: + def __init__(self, items=None): + self.items = [] if items is None else items + + def clone(self): + return self.__class__(items=self.items.copy()) + + +class History: + def __init__(self, x, size): + self.history = torch.zeros(size, *x.shape, device=x.device, dtype=x.dtype) + self.size = size + self.pos = 0 + self.last = None + + def __len__(self): + return min(self.pos, self.size) + + def __getitem__(self, k): + idx = (self.pos + k if k < 0 else self.pos + -self.size + k) % self.size + # print(f"\nFETCH {k}: pos={self.pos}, size={self.size}, at={idx}") + return self.history[idx] + + def push(self, val): + # print(f"\nPUSH {self.pos % self.size}: pos={self.pos}, size={self.size}") + self.last = self.pos % self.size + self.history[self.last] = val + self.pos += 1 + + def reset(self): + self.pos = 0 + self.last = None + + +class SamplerState: + def __init__( + self, + model, + sigmas, + idx, + s_in, + dhist, + xhist, + extra_args, + *, + noise_sampler, + callback=None, + denoised=None, + model_call_cache=None, + eta=1.0, + reta=1.0, + s_noise=1.0, + ): + self.model_ = model + self.dhist = dhist + self.xhist = xhist + self.extra_args = extra_args + self.s_in = s_in + self.eta = eta + self.reta = reta + self.s_noise = s_noise + self.sigmas = sigmas + self.denoised = denoised + self.callback_ = callback + self.noise_sampler = noise_sampler + self.model_call_cache = model_call_cache + self.update(idx) + + def update(self, idx=None): + idx = self.idx if idx is None else idx + 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() + self.sigma_down, self.sigma_up = get_ancestral_step( + self.sigma, self.sigma_next, eta=self.eta + ) + self.sigma_down_reversible, self.sigma_up_reversible = get_ancestral_step( + self.sigma, self.sigma_next, eta=self.reta + ) + + def model(self, x, sigma, *, model_call_idx=0, **kwargs): + mcc = self.model_call_cache + if mcc is None or model_call_idx >= mcc.size: + return self.model_(x, sigma * self.s_in, **self.extra_args, **kwargs) + if model_call_idx < mcc.pos: + print("CACHED MODEL CALL", model_call_idx) + return mcc.history[model_call_idx] + result = self.model_(x, sigma * self.s_in, **self.extra_args, **kwargs) + mcc.push(result) + print("CACHING MODEL CALL", model_call_idx, mcc.size, mcc.pos) + return result + + def get_ancestral_step(self, eta=1.0): + return get_ancestral_step(self.sigma, self.sigma_next, eta=eta) + + def clone_edit(self, **kwargs): + obj = self.__class__.__new__(self.__class__) + for k in ( + "model_", + "dhist", + "xhist", + "model_call_cache", + "extra_args", + "s_in", + "eta", + "reta", + "s_noise", + "sigmas", + "denoised", + "callback_", + "noise_sampler", + "idx", + "sigma", + "sigma_next", + "sigma_prev", + "sigma_down", + "sigma_up", + "sigma_down_reversible", + "sigma_up_reversible", + ): + setattr(obj, k, kwargs[k] if k in kwargs else getattr(self, k)) + obj.update() + return obj + + 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], + } + ) diff --git a/py/utils.py b/py/utils.py new file mode 100644 index 0000000..33f29a8 --- /dev/null +++ b/py/utils.py @@ -0,0 +1,22 @@ +import math +import torch + + +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 + fun = torch.gt if desc else torch.lt + first_unsorted = fun(tensor[1:], tensor[:-1]).nonzero().flatten()[:1].add_(1) + return None if not len(first_unsorted) else first_unsorted.item()