From e7db5e921390f07268df1be8a3fdea5504acef03 Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 7 Jul 2024 06:58:55 -0600 Subject: [PATCH] Sync changes --- README.md | 16 +++- py/nodes.py | 2 +- py/noise.py | 59 ++++++++++-- py/substep_samplers.py | 201 ++++++++++++----------------------------- 4 files changed, 121 insertions(+), 157 deletions(-) diff --git a/README.md b/README.md index 678fe20..3dcd838 100644 --- a/README.md +++ b/README.md @@ -107,11 +107,18 @@ noise: # Batch size, 0 disables. size: 0 - # Reference mode, one of: + # Reference mode, values can be one of: # x: Uses the current latent as a reference. # noise: Uses the current noise as a reference (x - denoised) - # denoised: Uses the model image prediction as a reference. - mode: x + # denoised: Uses the model image prediction as a reference (factors in positive and negative prompts). + # uncond: The model unconditional prediction (negative prompt) + # cond: The model conditional prediction (positive prompt) + # Advanced feature: Additionally you may enter a string of operations in the format: + # "x - denoised * 2 + cond" (just an example, not a recommended setting) + # Possible operations: + - / * min max add sub div mul + # Note: Each value and operation must be space delimited (i.e. "x-1" will not work). + # Also normal operator precedence does not apply here. + ref: x # Batching mode, one of: # batch: Matches vs batches. Immiscible mode is disabled if size < 2 @@ -133,6 +140,9 @@ noise: # You get (immiscible_noise * strength) + ((1.0 - strength) * normal_noise) - LERP. strength: 1.0 + # See: https://docs.scipy.org/doc/scipy/reference/generated/scipy.optimize.linear_sum_assignment.html#scipy.optimize.linear_sum_assignment + maximize: false + # Model calls can be cached. This is very experimental: I don't recommend using it # unless you know what you're doing. diff --git a/py/nodes.py b/py/nodes.py index 323a86a..9b360dd 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -8,7 +8,7 @@ from .substep_samplers import STEP_SAMPLERS from .substep_merging import MERGE_SUBSTEPS_CLASSES DEFAULT_YAML_PARAMS = """\ -# Enter parameters here in JSON or YAML format +# JSON or YAML parameters s_noise: 1.0 eta: 1.0 """ diff --git a/py/noise.py b/py/noise.py index 09b877f..1733161 100644 --- a/py/noise.py +++ b/py/noise.py @@ -11,8 +11,16 @@ from .utils import scale_noise, fallback class ImmiscibleNoise( collections.namedtuple( "ImmiscibleConfig", - ("size", "mode", "batching", "scale_ref", "normalize_ref", "strength"), - defaults=(0, "x", "channels", 1.0, False, 1.0), + ( + "size", + "ref", + "batching", + "scale_ref", + "normalize_ref", + "strength", + "maximize", + ), + defaults=(0, "x", "channels", 1.0, False, 1.0, False), ) ): def __call__(self, noise_sampler, x_ref, refs=None): @@ -32,9 +40,9 @@ class ImmiscibleNoise( ) # LERP def get_ref(self, x_ref, refs=None): - ref = fallback(refs, {}).get(self.mode, x_ref) + ref = self.custom_ref({"x": x_ref} | fallback(refs, {})) ref = scale_noise( - ref.clone(), + ref, self.scale_ref, normalized=self.normalize_ref not in (False, None), normalize_dims=self.normalize_ref @@ -43,6 +51,41 @@ class ImmiscibleNoise( ) return ref + def custom_ref(self, refs): + ops = self.ref.split(None) + ops.reverse() + result = refs.get(ops.pop()) + if result is None: + raise ValueError("Bad custom refs: must start with a value") + result = result.clone() + while ops: + if len(ops) < 2: + raise ValueError("Bad custom refs: too short") + op, valkey = ops.pop(), ops.pop() + if valkey[0] in ("+", "-") or valkey[0].isdigit(): + val = float(valkey) + else: + val = refs.get(valkey) + if val is None: + raise ValueError(f"Value {valkey} not found for op {op}") + if op in ("+", "add"): + result += val + elif op in ("-", "sub"): + result -= val + elif op in ("*", "mul"): + result *= val + elif op in ("/", "div"): + result /= min(val, 1e-06) + elif op in ("min", "max"): + if not isinstance(val, torch.Tensor): + val = result.new_full((1,), val) + result = result.minimum(val) if op == "min" else result.maximum(val) + elif op == "norm": + result = scale_noise(result, val, normalized=True) + if result is None: + raise ValueError("Bad custom reference") + return result + def batch(self, noise, ref): if self.batching == "batch": return noise, ref @@ -71,8 +114,7 @@ class ImmiscibleNoise( # Based on implementation from https://github.com/kohya-ss/sd-scripts/pull/1395 # Idea from https://github.com/Clybius - @staticmethod - def choose(noise, latents): + def choose(self, noise, latents): # "Immiscible Diffusion: Accelerating Diffusion Training with Noise Assignment" (2024) Li et al. arxiv.org/abs/2406.12303 # Minimize latent-noise pairs over a batch n = noise.shape[0] @@ -82,8 +124,8 @@ class ImmiscibleNoise( ) dist = (latents_expanded - noise_expanded) ** 2 dist = dist.mean(list(range(2, dist.dim()))).cpu() - assign_mat = scipy.optimize.linear_sum_assignment(dist) - print("IMM IDX", assign_mat[1]) + assign_mat = scipy.optimize.linear_sum_assignment(dist, maximize=self.maximize) + # print("IMM IDX", assign_mat[1]) return noise[assign_mat[1]] @@ -121,7 +163,6 @@ class NoiseSamplerCache: self.scale = float(scale) self.normalize_dims = tuple(int(v) for v in normalize_dims) self.immiscible = ImmiscibleNoise(**fallback(immiscible, {})) - print("IMM", self.immiscible) self.update_x(x) if set_seed: random.seed(seed) diff --git a/py/substep_samplers.py b/py/substep_samplers.py index 8c9f58c..e223fe3 100644 --- a/py/substep_samplers.py +++ b/py/substep_samplers.py @@ -59,16 +59,33 @@ class SamplerResult: else: self.denoised = self.noise_pred = None _ = self.extract_pred(ss) + self.denoised_uncond = ss.hcur.denoised_uncond + self.denoised_cond = ss.hcur.denoised_cond def get_noise(self, scaled=True): - if self.sigma_next == 0: + if self.sigma_next == 0 or self.noise_scale == 0: return torch.zeros_like(self.x_) + refs = { + k: getattr(self, ak) + for k, ak in ( + ("x", "x"), + ("noise", "noise_pred"), + ("denoised", "denoised"), + ("uncond", "denoised_uncond"), + ("cond", "denoised_cond"), + ("sigma", "sigma"), + ("sigma_next", "sigma_next"), + ("sigma_down", "sigma_down"), + ("sigma_up", "sigma_up"), + ) + if getattr(self, ak) is not None + } return self.noise_sampler( self.sigma, self.sigma_next, out_hw=self.x.shape[-2:], - x_ref=self.noise_pred, - refs={"noise": self.noise_pred, "denoised": self.denoised}, + x_ref=self.x, + refs=refs, ).mul_(self.noise_scale if scaled else 1.0) def extract_pred(self, ss): @@ -90,14 +107,9 @@ class SamplerResult: def noise_x(self, x=None, scale=1.0): x = fallback(x, self.x) - # if x is None: - # x = self.x - # else: - # self.x = x if self.sigma_next == 0 or self.noise_scale == 0: return x x = x + self.get_noise() * scale - # self.x = x + self.get_noise() * scale return x def clone(self): @@ -137,6 +149,7 @@ class SingleStepSampler(CFGPPStepMixin): self_noise = 0 model_calls = 0 ancestralize = False + sample_sigma_zero = False def __init__( self, @@ -161,14 +174,19 @@ class SingleStepSampler(CFGPPStepMixin): self.substeps = substeps def __call__(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) - for sr in self.step(x, ss): - if sr.final and self.ancestralize: - sr = self.ancestralize_result(ss, sr) - if sr.final: - return (yield sr) - yield sr + if not self.sample_sigma_zero and ss.sigma_next == 0: + return (yield from self.denoised_result(ss)) + next_x = None + sg = self.step(x, ss) + with contextlib.suppress(StopIteration): + while True: + sr = sg.send(next_x) + if sr.final: + if self.ancestralize: + sr = self.ancestralize_result(ss, sr) + return (yield sr) + next_x = sr.x + yield sr def step(self, x, ss): raise NotImplementedError @@ -393,7 +411,7 @@ class DPMPP2MSDEStep(HistorySingleStepSampler): + (-h - eta_h).expm1().neg() * denoised ) noise_strength = ss.sigma_next * (-2 * eta_h).expm1().neg().sqrt() - if ss.sigma_next == 0 or self.available_history(ss) == 0: + if self.available_history(ss) == 0: return (yield from self.result(ss, x, noise_strength)) h_last = (-ss.sigma.log()) - (-ss.sigma_prev.log()) r = h_last / h @@ -416,11 +434,7 @@ class DPMPP3MSDEStep(HistorySingleStepSampler): default_history_limit, max_history = 2, 2 def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.denoised_result(ss)) denoised = ss.denoised - # 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) @@ -462,8 +476,6 @@ class ReversibleHeunStep(ReversibleSingleStepSampler): allow_cfgpp = True def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) reta, reversible_scale = self.get_reversible_cfg(ss) sigma_down_reversible, _sigma_up_reversible = ss.get_ancestral_step(reta) @@ -500,8 +512,6 @@ class ReversibleHeun1SStep(ReversibleSingleStepSampler): allow_cfgpp = True def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) if self.available_history(ss) < 1: return (yield from ReversibleHeunStep.step(self, x, ss)) s = ss.sigma @@ -629,8 +639,6 @@ class RESStep(SingleStepSampler): self.c2 = res_c2 def step(self, x, ss): - if ss.sigma_next == 0: - 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 @@ -661,8 +669,6 @@ class TrapezoidalStep(SingleStepSampler): allow_alt_cfgpp = True def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) sigma_down, sigma_up = ss.get_ancestral_step(self.get_dyn_eta(ss)) # Calculate the derivative using the model @@ -689,9 +695,6 @@ class TrapezoidalCycleStep(CycleSingleStepSampler): allow_alt_cfgpp = True def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.denoised_result(ss)) - # Calculate the derivative using the model d_i = self.to_d(ss.hcur) @@ -724,8 +727,6 @@ class BogackiStep(ReversibleSingleStepSampler): self.reversible_scale = 0 def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) s = ss.sigma sd, su = ss.get_ancestral_step(self.get_dyn_eta(ss)) reta, reversible_scale = self.get_reversible_cfg(ss) @@ -765,8 +766,6 @@ class RK4Step(SingleStepSampler): allow_alt_cfgpp = True def step(self, x, ss): - if ss.sigma_next == 0: - 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 @@ -811,8 +810,6 @@ 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 :] @@ -1015,8 +1012,6 @@ class DPMPP2SStep(SingleStepSampler, DPMPPStepMixin): model_calls = 1 def step(self, x, ss): - if ss.sigma_next == 0: - 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) @@ -1040,8 +1035,6 @@ class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin): self.r = r def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) t_fn, sigma_fn = self.t_fn, self.sigma_fn r, eta = self.r, self.get_dyn_eta(ss) # DPM-Solver++ @@ -1078,8 +1071,6 @@ class TTMJVPStep(SingleStepSampler): self.alternate_phi_2_calc = alternate_phi_2_calc def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.denoised_result(ss)) eta = self.get_dyn_eta(ss) sigma, sigma_next = ss.sigma, ss.sigma_next # 2nd order truncated Taylor method @@ -1123,8 +1114,6 @@ class IPNDMStep(HistorySingleStepSampler): ) def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) order = self.available_history(ss) + 1 if order > 1: hd = tuple(self.to_d(ss.hist[-hidx]) for hidx in range(order, 1, -1)) @@ -1145,8 +1134,6 @@ class IPNDMVStep(HistorySingleStepSampler): allow_alt_cfgpp = True def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) dt = ss.dt d = self.to_d(ss.hcur) order = self.available_history(ss) + 1 @@ -1250,8 +1237,6 @@ class DEISStep(HistorySingleStepSampler): return self.deis_coeffs def step(self, x, ss): - if ss.sigma_next == 0: - return (yield from self.euler_step(x, ss)) dt = ss.dt d = self.to_d(ss.hcur) order = self.available_history(ss) + 1 @@ -1280,7 +1265,7 @@ class HeunPP2Step(SingleStepSampler): steps_remain = max(0, len(ss.sigmas) - (ss.idx + 2)) order = min(self.max_order, steps_remain + 1) sn = ss.sigma_next - if order == 1 or sn == 0: + if order == 1: return (yield from self.euler_step(x, ss)) d = self.to_d(ss.hcur) dt = ss.dt @@ -1306,6 +1291,7 @@ class HeunPP2Step(SingleStepSampler): class DESolverStep(SingleStepSampler, MinSigmaStepMixin): de_default_solver = None + sample_sigma_zero = True def __init__( self, @@ -1373,16 +1359,6 @@ class TDEStep(DESolverStep): def step(self, x, ss): s, sn, sigma_down, sigma_up = self.de_get_step(ss, x) - # eta = self.get_dyn_eta(ss) - # s, sn = ss.sigma, ss.sigma_next - # if s <= self.de_min_sigma: - # return (yield from self.euler_step(x, ss)) - # sn = self.adjust_step(sn, self.de_min_sigma) - # sigma_down, sigma_up = ss.get_ancestral_step(eta, sigma_next=sn) - # if self.de_fixup_hack != 0: - # sigma_down = (sigma_down - (s - sigma_down) * self.de_fixup_hack).clamp( - # min=0 - # ) delta = (s - sigma_down).item() mcc = 0 bidx = 0 @@ -1453,21 +1429,16 @@ class TDEStep(DESolverStep): yield from self.result(ss, result, sigma_up, sigma_down=sigma_down) -class TODEStep(SingleStepSampler, MinSigmaStepMixin): +class TODEStep(DESolverStep): name = "tode" model_calls = 2 allow_alt_cfgpp = True + de_default_solver = "dopri5" def __init__( self, *args, - de_solver="dopri5", - de_max_nfe=100, - de_rtol=-1.5, - de_atol=-3.5, - de_fixup_hack=0.025, de_initial_step=0.25, - de_min_sigma=0.0292, de_compile=False, de_ctl_pcoeff=0.3, de_ctl_icoeff=0.9, @@ -1479,31 +1450,21 @@ class TODEStep(SingleStepSampler, MinSigmaStepMixin): "TODE sampler requires torchode installed in venv. Example: pip install torchode" ) super().__init__(*args, **kwargs) - self.de_solver_name = de_solver - self.de_solver_method = tode.interface.METHODS[de_solver] - self.de_max_nfe = de_max_nfe - self.de_rtol = 10**de_rtol - self.de_atol = 10**de_atol + self.de_solver_method = tode.interface.METHODS[self.de_solver_name] self.de_ctl_pcoeff = de_ctl_pcoeff self.de_ctl_icoeff = de_ctl_icoeff self.de_ctl_dcoeff = de_ctl_dcoeff - self.de_fixup_hack = de_fixup_hack self.de_compile = de_compile - self.de_min_sigma = de_min_sigma if de_min_sigma is not None else 0.0 self.de_initial_step = de_initial_step - def step(self, x, ss): - eta = self.get_dyn_eta(ss) - s, sn = ss.sigma, ss.sigma_next - if s <= self.de_min_sigma: - return (yield from self.euler_step(x, ss)) - sn = self.adjust_step(sn, self.de_min_sigma) - sigma_down, sigma_up = ss.get_ancestral_step(eta, sigma_next=sn) - if self.de_fixup_hack != 0: - sigma_down = (sigma_down - (s - sigma_down) * self.de_fixup_hack).clamp( - min=0 + def check_solver_support(self): + if not HAVE_TODE: + raise RuntimeError( + "TODE sampler requires torchode installed in venv. Example: pip install torchode" ) + def step(self, x, ss): + s, sn, sigma_down, sigma_up = self.de_get_step(ss, x) delta = (ss.sigma - sigma_down).item() mcc = 0 pbar = None @@ -1579,64 +1540,44 @@ class TODEStep(SingleStepSampler, MinSigmaStepMixin): yield from self.result(ss, result, sigma_up, sigma_down=sigma_down) -class TSDEStep(SingleStepSampler, MinSigmaStepMixin): +class TSDEStep(DESolverStep): name = "tsde" model_calls = 2 allow_alt_cfgpp = True + de_default_solver = "reversible_heun" def __init__( self, *args, - de_solver="euler", - de_max_nfe=100, - de_rtol=-1.5, - de_atol=-3.5, - de_fixup_hack=0.025, de_initial_step=0.25, - de_min_sigma=0.0292, de_split=1, de_adaptive=False, de_noise_type="scalar", - de_sde_type="ito", + de_sde_type="stratonovich", + de_levy_area_approx="none", de_noise_channels=1, de_g_multiplier=0.05, de_g_reverse_time=True, de_g_derp_mode=False, **kwargs, ): - if not HAVE_TODE: - raise RuntimeError( - "TODE sampler requires torchode installed in venv. Example: pip install torchode" - ) super().__init__(*args, **kwargs) - self.de_solver_name = de_solver - self.de_max_nfe = de_max_nfe - self.de_rtol = 10**de_rtol - self.de_atol = 10**de_atol - self.de_fixup_hack = de_fixup_hack - self.de_min_sigma = de_min_sigma if de_min_sigma is not None else 0.0 self.de_initial_step = de_initial_step self.de_adaptive = de_adaptive self.de_split = de_split self.de_noise_type = de_noise_type self.de_sde_type = de_sde_type + self.de_levy_area_approx = de_levy_area_approx self.de_g_multiplier = de_g_multiplier self.de_noise_channels = de_noise_channels self.de_g_reverse_time = de_g_reverse_time self.de_g_derp_mode = de_g_derp_mode - def step(self, x, ss): - eta = self.get_dyn_eta(ss) - s, sn = ss.sigma, ss.sigma_next - if s <= self.de_min_sigma: - return (yield from self.euler_step(x, ss)) - sn = self.adjust_step(sn, self.de_min_sigma) - sigma_down, sigma_up = ss.get_ancestral_step(eta, sigma_next=sn) - if self.de_fixup_hack != 0: - sigma_down = (sigma_down - (s - sigma_down) * self.de_fixup_hack).clamp( - min=0 - ) + def check_solver_support(self): + pass + def step(self, x, ss): + s, sn, sigma_down, sigma_up = self.de_get_step(ss, x) delta = (ss.sigma - sigma_down).item() mcc = 0 pbar = None @@ -1650,7 +1591,6 @@ class TSDEStep(SingleStepSampler, MinSigmaStepMixin): @torch.no_grad() def f(self, t_rev, y_flat): nonlocal mcc - # t = t_rev.abs() t = s - (t_rev - sigma_down) # print(f"\nf at t_rev={t_rev}, t={t} :: {y_flat.shape}") if torch.all(t <= 1e-05).item(): @@ -1683,25 +1623,17 @@ class TSDEStep(SingleStepSampler, MinSigmaStepMixin): pct = t / (s - sigma_down) if outer_self.de_g_reverse_time: pct = 1.0 - pct - # print("\n>>", t, t_rev, "--", t * outer_self.de_g_multiplier) multiplier = outer_self.de_g_multiplier if outer_self.de_g_derp_mode and mcc % 2 == 0: multiplier *= -1 val = t * pct * multiplier - # out = ( - # (val * y_flat) - # .repeat(1, 1, outer_self.de_noise_channels) - # .view(*y_flat.shape, outer_self.de_noise_channels) - # ) if self.noise_type == "diagonal": out = val.repeat(*y_flat.shape) elif self.noise_type == "scalar": out = val.repeat(*y_flat.shape, 1) else: out = val.repeat(*y_flat.shape, outer_self.de_noise_channels) - print("\nOUT", val, out.shape) return out - # return val.repeat(*y_flat.shape, outer_self.de_noise_channels) t = torch.stack((sigma_down, s)).to(torch.float) @@ -1712,18 +1644,6 @@ class TSDEStep(SingleStepSampler, MinSigmaStepMixin): dt0 = ( delta * self.de_initial_step if self.de_adaptive else delta / self.de_split ) - # bm = torchsde.BrownianTree( - # t0=s, - # # t1=sigma_down, - # entropy=ss.noise.seed, - # w0=torch.zeros( - # b, - # self.de_noise_channels, - # dtype=x.dtype, - # device=x.device, - # ), - # # levy_area_approximation="space-time", - # ) sde = SDE() y_flat = x.flatten(start_dim=1) if sde.noise_type == "diagonal": @@ -1738,9 +1658,7 @@ class TSDEStep(SingleStepSampler, MinSigmaStepMixin): t0=-s, t1=s, entropy=ss.noise.seed, - levy_area_approximation="space-time" - if self.de_solver_name != "log_ode" - else "davie", + levy_area_approximation=self.de_levy_area_approx, tol=1e-06, size=bm_size, ) @@ -1785,9 +1703,7 @@ class HeunStep(ReversibleSingleStepSampler): return (dtr**2 * (d_to - d_from) / 4) * self.reversible_scale def step(self, x, ss): - s, sn = ss.sigma, ss.sigma_next - if sn <= 0: - return (yield from self.euler_step(x, ss)) + s = ss.sigma sd, su = ss.get_ancestral_step(self.get_dyn_eta(ss)) dt = sd - s hcur = ss.hcur @@ -1808,10 +1724,7 @@ class Heun1SStep(HeunStep): default_history_limit, max_history = 1, 1 def step(self, x, ss): - sn = ss.sigma_next s = ss.sigma - if sn <= 0: - return (yield from self.euler_step(x, ss)) if self.available_history(ss) == 0: return (yield from super().step(x, ss)) hcur, hprev = ss.hcur, ss.hprev