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:
blepping
2024-12-07 13:48:27 -07:00
parent 044490f2e5
commit dac98129ca
9 changed files with 134 additions and 31 deletions
+11
View File
@@ -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(),
+6
View File
@@ -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
View File
@@ -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,
+30 -3
View File
@@ -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
+28 -4
View File
@@ -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)),
+7 -1
View File
@@ -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
+1 -1
View File
@@ -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__(
+26 -16
View File
@@ -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
View File
@@ -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