Sync changes
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user