From dac98129caaa88e511d776c9defbb998c8b97a7a Mon Sep 17 00:00:00 2001 From: blepping Date: Sat, 7 Dec 2024 13:48:27 -0700 Subject: [PATCH] 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 --- py/expression_handlers.py | 11 +++++++++ py/filtering.py | 6 +++++ py/noise.py | 18 ++++++++++---- py/step_samplers/base.py | 33 ++++++++++++++++++++++--- py/step_samplers/blep.py | 32 +++++++++++++++++++++--- py/step_samplers/builtins.py | 8 +++++- py/step_samplers/clybius.py | 2 +- py/step_samplers/extraltodeus.py | 42 ++++++++++++++++++++------------ py/substep_merging.py | 13 +++++++++- 9 files changed, 134 insertions(+), 31 deletions(-) diff --git a/py/expression_handlers.py b/py/expression_handlers.py index 4ce175e..a95a693 100644 --- a/py/expression_handlers.py +++ b/py/expression_handlers.py @@ -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(), diff --git a/py/filtering.py b/py/filtering.py index f147713..0c8bab9 100644 --- a/py/filtering.py +++ b/py/filtering.py @@ -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"" + class SimpleFilter(Filter): name = "simple" diff --git a/py/noise.py b/py/noise.py index f88d25d..b5d6a83 100644 --- a/py/noise.py +++ b/py/noise.py @@ -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, diff --git a/py/step_samplers/base.py b/py/step_samplers/base.py index 67f247c..5d8a5bc 100644 --- a/py/step_samplers/base.py +++ b/py/step_samplers/base.py @@ -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 diff --git a/py/step_samplers/blep.py b/py/step_samplers/blep.py index 9f7718e..d0e9571 100644 --- a/py/step_samplers/blep.py +++ b/py/step_samplers/blep.py @@ -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)), diff --git a/py/step_samplers/builtins.py b/py/step_samplers/builtins.py index 3267b6e..55489b3 100644 --- a/py/step_samplers/builtins.py +++ b/py/step_samplers/builtins.py @@ -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 diff --git a/py/step_samplers/clybius.py b/py/step_samplers/clybius.py index e981ce5..5fa5816 100644 --- a/py/step_samplers/clybius.py +++ b/py/step_samplers/clybius.py @@ -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__( diff --git a/py/step_samplers/extraltodeus.py b/py/step_samplers/extraltodeus.py index 7375bdd..b28f519 100644 --- a/py/step_samplers/extraltodeus.py +++ b/py/step_samplers/extraltodeus.py @@ -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: diff --git a/py/substep_merging.py b/py/substep_merging.py index 7158fed..04109de 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -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