Sync changes

This commit is contained in:
blepping
2024-07-07 06:58:55 -06:00
parent 361da0f32b
commit e7db5e9213
4 changed files with 121 additions and 157 deletions
+13 -3
View File
@@ -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.
+1 -1
View File
@@ -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
"""
+50 -9
View File
@@ -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)
+57 -144
View File
@@ -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