Stage 1
This commit is contained in:
@@ -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"]
|
||||
|
||||
+200
-73
@@ -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",
|
||||
)
|
||||
|
||||
+23
-37
@@ -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
|
||||
|
||||
+248
-122
@@ -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,
|
||||
}
|
||||
|
||||
+268
-49
@@ -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"<SS({self.name}): s_noise={self.s_noise}, eta={self.eta}>"
|
||||
@@ -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 = {
|
||||
|
||||
+222
-13
@@ -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],
|
||||
})
|
||||
|
||||
+33
-7
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user