Fix noise sampler caching not considering immiscible settings
Allow specifying alt custom noise/immiscible settings for samplers that internally add noise AFS support at the merge sampler level Allow the Weoon internal step to use ETA Add t_copysign expression function
This commit is contained in:
@@ -123,6 +123,16 @@ class FlipHandler(NormHandler):
|
||||
return result
|
||||
|
||||
|
||||
class CopySignHandler(NormHandler):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("tensor"),
|
||||
expr.Arg.tensor("other"),
|
||||
)
|
||||
|
||||
def handle(self, obj, getter):
|
||||
return torch.copysign(*self.safe_get_all(obj, getter))
|
||||
|
||||
|
||||
class BlendHandler(NormHandler):
|
||||
input_validators = (
|
||||
expr.Arg.tensor("tensor1"),
|
||||
@@ -765,6 +775,7 @@ TENSOR_OP_HANDLERS = {
|
||||
"t_blend": BlendHandler(),
|
||||
"t_roll": RollHandler(),
|
||||
"t_flip": FlipHandler(),
|
||||
"t_copysign": CopySignHandler(),
|
||||
"t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(),
|
||||
"t_scale": ScaleHandler(),
|
||||
"t_noise": NoiseHandler(),
|
||||
|
||||
@@ -246,6 +246,12 @@ class Filter:
|
||||
refs = drefs if refs is None else refs | drefs
|
||||
return ops.eval(FILTER_HANDLERS.clone(constants=refs, variables={}))
|
||||
|
||||
def __str__(self):
|
||||
prettyvals = ", ".join(
|
||||
f"{k}={getattr(self, k, None)!s}" for k in self.default_options.keys()
|
||||
)
|
||||
return f"<Filter({self.name}): {prettyvals}>"
|
||||
|
||||
|
||||
class SimpleFilter(Filter):
|
||||
name = "simple"
|
||||
|
||||
+13
-5
@@ -182,7 +182,9 @@ class NoiseSamplerCache:
|
||||
sigmas=None,
|
||||
):
|
||||
size = min(size, self.batch_size)
|
||||
cache_key = (nsobj, size)
|
||||
if immiscible is None:
|
||||
immiscible = self.immiscible
|
||||
cache_key = (nsobj, size, hash(immiscible))
|
||||
if self.caching:
|
||||
noise_sampler = self.cache.get(cache_key)
|
||||
if noise_sampler:
|
||||
@@ -192,8 +194,16 @@ class NoiseSamplerCache:
|
||||
curr_x = self.mega_x[: self.x.shape[0] * size, ...]
|
||||
if nsobj is None:
|
||||
|
||||
def ns(_s, _sn, *_unused, **_unusedkwargs):
|
||||
return torch.randn_like(curr_x)
|
||||
def ns(*_unused, **_unusedkwargs):
|
||||
noise = torch.randn(
|
||||
curr_x.shape,
|
||||
dtype=curr_x.dtype,
|
||||
layout=curr_x.layout,
|
||||
device="cpu" if self.cpu_noise else curr_x.device,
|
||||
)
|
||||
if noise.device != curr_x.device:
|
||||
return noise.to(curr_x.device)
|
||||
return noise
|
||||
|
||||
else:
|
||||
if sigmas is not None:
|
||||
@@ -212,8 +222,6 @@ class NoiseSamplerCache:
|
||||
orig_h, orig_w = self.x.shape[-2:]
|
||||
remain = 0
|
||||
noise = None
|
||||
if immiscible is None:
|
||||
immiscible = self.immiscible
|
||||
|
||||
def noise_sampler_(
|
||||
curr_sigma,
|
||||
|
||||
@@ -98,12 +98,12 @@ class SamplerResult:
|
||||
x = fallback(x, self.x)
|
||||
if self.sigma_next == 0 or self.noise_scale == 0:
|
||||
return x
|
||||
noise = self.get_noise(ss=ss) * scale
|
||||
noise = self.get_noise(ss=ss).mul_(scale)
|
||||
if not self.is_rectified_flow:
|
||||
return x + noise
|
||||
return noise.add_(x)
|
||||
x_coeff = (1 - self.sigma_next) / (1 - self.sigma_down)
|
||||
# print(f"\nRF noise: {x_coeff}")
|
||||
return x_coeff * x + noise
|
||||
return noise.add_(x_coeff * x)
|
||||
|
||||
def clone(self):
|
||||
obj = self.__new__(self.__class__)
|
||||
@@ -139,6 +139,7 @@ class SingleStepSampler:
|
||||
allow_cfgpp = False
|
||||
allow_alt_cfgpp = False
|
||||
afs_end_step = -1
|
||||
uses_alt_noise = False
|
||||
|
||||
default_eta = 1.0
|
||||
|
||||
@@ -186,6 +187,15 @@ class SingleStepSampler:
|
||||
self.custom_noise = self.options.get("custom_noise")
|
||||
if isinstance(self.custom_noise, str):
|
||||
self.custom_noise = self.options.get(f"custom_noise_{self.custom_noise}")
|
||||
if not self.uses_alt_noise:
|
||||
return
|
||||
self.alt_custom_noise = self.options.get("custom_noise_alt")
|
||||
alt_immiscible = self.options.get("alt_immiscible")
|
||||
self.alt_immiscible = (
|
||||
noise.ImmiscibleNoise(**alt_immiscible)
|
||||
if isinstance(alt_immiscible, dict)
|
||||
else alt_immiscible
|
||||
)
|
||||
|
||||
def __call__(self, x):
|
||||
ss = self.ss
|
||||
@@ -225,10 +235,27 @@ class SingleStepSampler:
|
||||
ss.sigma_next,
|
||||
immiscible=fallback(self.immiscible, ss.noise.immiscible),
|
||||
)
|
||||
if not self.uses_alt_noise:
|
||||
return
|
||||
if self.alt_custom_noise is None and self.alt_immiscible is None:
|
||||
self.alt_noise_sampler = self.noise_sampler
|
||||
return
|
||||
self.alt_noise_sampler = ss.noise.make_caching_noise_sampler(
|
||||
fallback(self.alt_custom_noise, self.custom_noise),
|
||||
1,
|
||||
ss.sigma,
|
||||
ss.sigma_next,
|
||||
immiscible=fallback(
|
||||
fallback(self.alt_immiscible, self.immiscible),
|
||||
ss.noise.immiscible,
|
||||
),
|
||||
)
|
||||
|
||||
def reset(self):
|
||||
self.ss = None
|
||||
self.noise_sampler = None
|
||||
if self.uses_alt_noise:
|
||||
self.alt_noise_sampler = None
|
||||
|
||||
def afs_step(self, x):
|
||||
sigma, sigma_next = self.ss.sigma, self.ss.sigma_next
|
||||
|
||||
@@ -245,13 +245,14 @@ class BASConfig(typing.NamedTuple):
|
||||
class BASStep(SingleStepSampler):
|
||||
name = "blep_bas"
|
||||
model_calls = -1
|
||||
uses_alt_noise = True
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.bas = BASConfig(**self.options.get("bas", {}))
|
||||
if self.bas.renoise_mode not in {"restart", "restart_noneta", "simple"}:
|
||||
bas = self.bas = BASConfig(**self.options.get("bas", {}))
|
||||
if bas.renoise_mode not in {"restart", "restart_noneta", "simple"}:
|
||||
raise ValueError("Bad BAS renoise mode")
|
||||
if self.bas.tostep_source not in {"dt", "sigma", "sigma_next"}:
|
||||
if bas.tostep_source not in {"dt", "sigma", "sigma_next"}:
|
||||
raise ValueError("Bad BAS tostep_source")
|
||||
blend_mode = self.options.get("blend_mode", "lerp").strip()
|
||||
self.blend = (
|
||||
@@ -333,7 +334,10 @@ class BASStep(SingleStepSampler):
|
||||
] = yield from self.result(
|
||||
x_new,
|
||||
noise_factor,
|
||||
sigma=bsigma,
|
||||
sigma_down=bsigma_down,
|
||||
s_noise=bas.s_noise,
|
||||
noise_sampler=self.alt_noise_sampler,
|
||||
final=False,
|
||||
)
|
||||
s_in = x.new_ones(expanded_batch)
|
||||
@@ -377,6 +381,9 @@ def blend_wavelets(a, b, *, factor_yl, factor_yh, blend_yl, blend_yh=None):
|
||||
class WeoonConfig(typing.NamedTuple):
|
||||
start_step: int = 0
|
||||
end_step: int = 9999
|
||||
eta: float = 0.0
|
||||
eta_retry_increment: float = 0.0
|
||||
s_noise: float = 1.0
|
||||
# One of dwt, dwt1d, dtcwt
|
||||
wavelet_mode: str = "dwt"
|
||||
padding: str = "periodization"
|
||||
@@ -403,6 +410,7 @@ class WeoonConfig(typing.NamedTuple):
|
||||
class WeoonStep(SingleStepSampler):
|
||||
name = "blep_weoon"
|
||||
model_calls = 1
|
||||
uses_alt_noise = True
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
if not HAVE_WAVELETS:
|
||||
@@ -481,8 +489,24 @@ class WeoonStep(SingleStepSampler):
|
||||
self.wavelet_inverse.to(x)
|
||||
dt = sigma_next - sigma
|
||||
wsigma_next = (sigma + dt * w.downstep_scale).clamp_(0)
|
||||
wratio = wsigma_next / sigma
|
||||
wsigma_down, wsigma_up = self.get_ancestral_step(
|
||||
w.eta,
|
||||
sigma=sigma,
|
||||
sigma_next=wsigma_next,
|
||||
retry_increment=w.eta_retry_increment,
|
||||
)
|
||||
wratio = wsigma_down / sigma
|
||||
x_down = self.blend(ss.denoised, x, wratio)
|
||||
if wsigma_up != 0:
|
||||
x_down = yield from self.result(
|
||||
x_down,
|
||||
wsigma_up,
|
||||
sigma_next=wsigma_next,
|
||||
sigma_down=wsigma_down,
|
||||
s_noise=w.s_noise,
|
||||
noise_sampler=self.alt_noise_sampler,
|
||||
final=False,
|
||||
)
|
||||
mr_down = self.call_model(x_down, wsigma_next, call_index=1)
|
||||
coeffs = scale_wavelets(
|
||||
self.wavelet_forward(self.maybe_flatten(ss.denoised)),
|
||||
|
||||
@@ -189,6 +189,7 @@ class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin):
|
||||
self_noise = 1
|
||||
model_calls = 1
|
||||
allow_alt_cfgpp = True # Implementation may not be correct.
|
||||
uses_alt_noise = True
|
||||
|
||||
def __init__(self, *args, r=1 / 2, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
@@ -214,7 +215,12 @@ class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin):
|
||||
)
|
||||
x_2 = (sigma_fn(s_) / sigma_fn(t)) * eff_x - (t - s_).expm1() * ss.denoised
|
||||
x_2 = yield from self.result(
|
||||
x_2, su, sigma=sigma_fn(t), sigma_next=sigma_fn(s), final=False
|
||||
x_2,
|
||||
su,
|
||||
sigma=sigma_fn(t),
|
||||
sigma_next=sigma_fn(s),
|
||||
noise_sampler=self.alt_noise_sampler,
|
||||
final=False,
|
||||
)
|
||||
denoised_2 = self.call_model(x_2, sigma_fn(s), call_index=1).denoised
|
||||
|
||||
|
||||
@@ -354,7 +354,7 @@ class RKDynamicStep(SingleStepSampler):
|
||||
|
||||
# Based on original implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
|
||||
class EulerDancingStep(SingleStepSampler):
|
||||
name = "euler_dancing"
|
||||
name = "clybius_euler_dancing"
|
||||
self_noise = 1
|
||||
|
||||
def __init__(
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# Samplers based on design from https://github.com/Extraltodeus/
|
||||
|
||||
import typing
|
||||
|
||||
import torch
|
||||
import tqdm
|
||||
|
||||
@@ -7,30 +9,36 @@ from .base import SingleStepSampler
|
||||
from . import registry
|
||||
|
||||
|
||||
class DistanceConfig(typing.NamedTuple):
|
||||
resample: int = 3
|
||||
resample_end: int = 1
|
||||
eta: float = 0.0
|
||||
s_noise: float = 1.0
|
||||
alt_cfgpp_scale: float = 0.0
|
||||
first_eta_step: int = 0
|
||||
last_eta_step: int = -1
|
||||
custom_noise_name: str = "alt"
|
||||
immiscible: dict | bool | None = None
|
||||
|
||||
|
||||
# Based on https://github.com/Extraltodeus/DistanceSampler
|
||||
class DistanceStep(SingleStepSampler):
|
||||
name = "distance"
|
||||
name = "extraltodeus_distance"
|
||||
allow_alt_cfgpp = True
|
||||
model_calls = -1
|
||||
uses_alt_noise = True
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
distance = self.options.get("distance", {})
|
||||
self.distance_resample = distance.get("resample", 3)
|
||||
self.distance_resample_end = distance.get("resample_end", 1)
|
||||
self.distance_resample_eta = distance.get("eta", self.eta)
|
||||
self.distance_resample_s_noise = distance.get("s_noise", self.s_noise)
|
||||
self.distance_alt_cfgpp_scale = distance.get("alt_cfgpp_scale", 0.0)
|
||||
self.distance_first_eta_resample_step = distance.get("first_eta_step", 0)
|
||||
self.distance_last_eta_resample_step = distance.get("last_eta_step", -1)
|
||||
self.distance = DistanceConfig(**self.options.get("distance", {}))
|
||||
|
||||
@property
|
||||
def require_uncond(self):
|
||||
return super().require_uncond or self.distance_alt_cfgpp_scale != 0
|
||||
return super().require_uncond or self.distance.alt_cfgpp_scale != 0
|
||||
|
||||
def distance_resample_steps(self):
|
||||
ss = self.ss
|
||||
resample, resample_end = self.distance_resample, self.distance_resample_end
|
||||
resample, resample_end = self.distance.resample, self.distance.resample_end
|
||||
if resample == -1:
|
||||
current_resample = min(10, (ss.sigmas.shape[0] - ss.idx) // 2)
|
||||
else:
|
||||
@@ -70,10 +78,11 @@ class DistanceStep(SingleStepSampler):
|
||||
resample_steps = self.distance_resample_steps()
|
||||
if resample_steps < 1:
|
||||
return (yield from self.euler_step(x))
|
||||
distance = self.distance
|
||||
ss = self.ss
|
||||
sigma_down, sigma_up = self.get_ancestral_step(self.get_dyn_eta())
|
||||
rsigma_down, rsigma_up = self.get_ancestral_step(eta=self.distance_resample_eta)
|
||||
rsigma_up *= self.distance_resample_s_noise
|
||||
rsigma_down, rsigma_up = self.get_ancestral_step(eta=distance.eta)
|
||||
rsigma_up *= distance.s_noise
|
||||
sigma, sigma_next = ss.sigma, ss.sigma_next
|
||||
zero_up = sigma * 0
|
||||
d = self.to_d(ss.hcur)
|
||||
@@ -81,8 +90,8 @@ class DistanceStep(SingleStepSampler):
|
||||
start_eta_idx, end_eta_idx = (
|
||||
max(0, resample_steps + v if v < 0 else v)
|
||||
for v in (
|
||||
self.distance_first_eta_resample_step,
|
||||
self.distance_last_eta_resample_step,
|
||||
distance.first_eta_step,
|
||||
distance.last_eta_step,
|
||||
)
|
||||
)
|
||||
dt = sigma_down - sigma
|
||||
@@ -103,11 +112,12 @@ class DistanceStep(SingleStepSampler):
|
||||
curr_sigma_up,
|
||||
sigma=sigma,
|
||||
sigma_down=curr_sigma_down,
|
||||
noise_sampler=self.alt_noise_sampler,
|
||||
final=False,
|
||||
)
|
||||
sr = self.call_model(x_new, sigma_next, call_index=re_step + 1)
|
||||
new_d = sr.to_d(
|
||||
sigma=curr_sigma_down, alt_cfgpp_scale=self.distance_alt_cfgpp_scale
|
||||
sigma=curr_sigma_down, alt_cfgpp_scale=distance.alt_cfgpp_scale
|
||||
)
|
||||
x_n.append(new_d)
|
||||
if re_step == 0:
|
||||
|
||||
+12
-1
@@ -39,6 +39,8 @@ class MergeSubstepsSampler:
|
||||
self.preview_mode = options.pop("preview_mode", "denoised")
|
||||
self.require_uncond = any(sampler.require_uncond for sampler in samplers)
|
||||
self.cfg_scale_override = options.pop("cfg_scale_override", None)
|
||||
self.afs_start_step = options.pop("afs_start_step", 0)
|
||||
self.afs_end_step = options.pop("afs_end_step", -1)
|
||||
self.options = options
|
||||
|
||||
def check_match(self, handlers: None | object, *, ss: None | object = None):
|
||||
@@ -74,9 +76,18 @@ class MergeSubstepsSampler:
|
||||
def __call__(self, x):
|
||||
orig_x = x
|
||||
x = self.step_input(x)
|
||||
x = self.step(x)
|
||||
if self.afs_start_step <= self.ss.step <= self.afs_end_step:
|
||||
x = self.afs_step(x)
|
||||
else:
|
||||
x = self.step(x)
|
||||
return self.step_output(x, orig_x=orig_x)
|
||||
|
||||
def afs_step(self, x):
|
||||
sigma, sigma_next = self.ss.sigma, self.ss.sigma_next
|
||||
afs_d = x / ((1 + sigma**2).sqrt())
|
||||
dt = sigma_next - sigma
|
||||
return x + afs_d * dt
|
||||
|
||||
def step(self, x):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
Reference in New Issue
Block a user