Initial implementation

This commit is contained in:
blepping
2024-05-30 07:47:43 -06:00
parent 360ba4842a
commit aa3e12ae34
10 changed files with 1328 additions and 2 deletions
+74 -2
View File
@@ -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!
+8
View File
@@ -0,0 +1,8 @@
from .py import nodes
NODE_CLASS_MAPPINGS = {
"ComposableSampler": nodes.ComposableSampler,
"ComposableStepSampler": nodes.ComposableStepSampler,
}
__all__ = ["NODE_CLASS_MAPPINGS"]
View File
+144
View File
@@ -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")
+127
View File
@@ -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,
)
+88
View File
@@ -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
+252
View File
@@ -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,
}
+469
View File
@@ -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"<SS({self.name}): s_noise={self.s_noise}, eta={self.eta}, kwargs={self.kwargs}>"
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",
)
+144
View File
@@ -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],
}
)
+22
View File
@@ -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()