All sorts of fun stuff

This commit is contained in:
blepping
2024-05-30 19:47:44 -06:00
parent 05dcec2985
commit be5e5d8933
6 changed files with 238 additions and 60 deletions
+21 -7
View File
@@ -7,6 +7,8 @@ Very unstable, experimental and mathematically unsound sampling for ComfyUI.
Current status: In flux, not suitable for general use.
*Note*: You will basically always have to tweak settings like `s_noise` to get a good result. If the generation looks smooth/undetailed increase `s_noise` somewhere. If it looks crunchy, super high contrast, etc then try reducing noise.
## Nodes
### ComposableSampler
@@ -14,16 +16,20 @@ Current status: In flux, not suitable for general use.
**Possible Parameters**
* `avgmerge_stretch`(`0.4`): Used for `average` and `sample` merge types. See below.
* `model_call_cache`(unset): Caches the result of model calls at n+1 (where `n` is the number of model evaluations per step). For example, Bogacki is three model calls per step: whether the first one runs is dependent on the merge strategy. After that, Bogacki calls the model two more times. If you set `model_call_cache` to `1` then the result of that second call will be cached and if you're running two Bogacki substeps then the second one will use the cached version. Massively accelerates inference (especially when using the `average` merge strategy) but is likely very unsound and inaccurate.
* `model_call_cache`(unset): Caches the result of model calls at n+1 (where `n` is the number of model evaluations per step). For example, Bogacki is three model calls per step: whether the first one runs is dependent on the merge strategy. After that, Bogacki calls the model two more times. If you set `model_call_cache` to `1` then the result of that second call will be cached and if you're running two Bogacki substeps then the second one will use the cached version. Massively accelerates inference (especially when using the `average` merge strategy) but is likely very unsound and inaccurate. Does not apply to the sampler call for the `sample` merge strategy.
* `model_call_cache_threshold`(`0`): Only has an effect when `model_call_cache` is set to `1` or higher. Disables using the first `n` cached model calls. For example, if set to `1` and using a sampler like Bogacki that calls the model two extra times, the first will never be cached.
#### Merging
When running multiple substeps per step, the results will combined based on the merge strategy. Possible strategies:
When running multiple substeps per step, the results will combined based on the merge strategy. Possible strategies (in order of least weird to most weird):
* `normal`: The model is called at least once per substep (and possibly additional times for higher order samplers). The result of each substep is noised and the next substep uses that result. Then all the results are averaged.
* `divide`: Creates a linear schedule between the current sigma and the next and runs the substeps in sequence. The model is called at least once per substep.
* `normal`: The model is called at least once per substep (and possibly additional times for higher order samplers). The result of each substep is noised and the next substep uses that result. Then all the results are averaged.
* `average`: The model is called once at the beginning of the step and substeps share that result (but it may be called additional times for higher order samplers). This means substeps for samplers like reversible Euler, Heun 1s, DPM++ 2m SDE are essentially free. May be theoretically very unsound and inaccurate, requires manual tweaking of settings like `s_noise`. Supports the parameter `avgmerge_stretch`(`0.4`) which basically rolls back the current sigma and adds some noise (otherwise running a substep is deterministic and there would be no point to running a sampler like Euler more than once).
* `sample`: Like `average` (and uses `avgmerge_stretch`) but instead of simply using the average, it does a sampler step toward it instead. You can plug in any substep sampler to the `merge_sampler_opt` input (if unconnected and the merge method is `sample` then Euler will be used).
* `sample`: Like `average` (and uses `avgmerge_stretch`) but instead of simply using the average, it does a sampler step toward that instead. You can plug in any substep sampler to the `merge_sampler_opt` input (if unconnected and the merge method is `sample` then Euler will be used). *Note*: Substeps in the attached sampler will be ignored.
* `sample_uncached`: Similar to `sample`, however it calls the model per substep instead of caching the result and sharing it. Aside from sampling toward the result, it works more like the `normal` merge strategy. Theoretically it should be better because it's taking less shortcuts but results seem worse.
When using `average` and `sample` merge strategies and with model call caching enabled you can get away with setting substeps super high. Running something like 100 substeps is actually quite practical and seems to work well.
### ComposableStepSampler
@@ -57,18 +63,26 @@ dyn_deta_mode: "deta"
* `dyn_deta_mode`(`lerp`): May be one of:
* `deta`: Scales `deta` based on the value from `dyn_deta_start/end`.
* `lerp`: Does the dance step according to `deta` and then LERPs the non-dance sample result with the dance sample result based on the scale calculated from `dyn_deta_start/end` (which is `1.0` if they are unset). For example, if the dance scale is `0.5` you will get 50% normal sampling, 50% dancing sampling.
* `lerp_alt`: Similar to `lerp` except it LERPs with the leap result instead of a normal Euler ancestral result.
#### RES
* `res_simple_phi`(`false`): Applies to RES. Uses a faster but possibly less accurate method for calculating phi. What does phi do? I haven't the foggiest!
* `res_c2`(`0.5`): Applies to RES. Solver partial step size, the default of `0.5` appears to use the midpoint. Setting it to a lower value might possibly be more accurate but slower?
* `res_simple_phi`(`false`): Uses a faster but possibly less accurate method for calculating phi. What does phi do? I haven't the foggiest!
* `res_c2`(`0.5`): Solver partial step size, the default of `0.5` appears to use the midpoint. Setting it to a lower value might possibly be more accurate but slower?
#### TTM JVP
`alterate_phi_2_calc`(`true`): Supposedly works better than disabled when ETA isn't 0. I didn't notice a difference.
**Note**: TTM is a weird sampler. If you're using model caching you must make sure the entries TTM uses are populated first (by having before any other samplers that call the model multiple times). It may also not work with some other model patches and upscale methods.
## Credits
I can move code around but sampling math and creating samplers is far beyond my ability. I didn't write any of the original samplers:
* Euler, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation.
* Euler, DPMPP SDE, DPMPP 2S, DPM++ 2m, 2m SDE and 3m SDE samplers based on ComfyUI's implementation.
* Reversible Heun, Reversible Heun 1s, RES, Trapezoidal, Bogacki, Reversible Bogacki, RK4 and Euler Dancing samplers based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
* TTM JVP sampler based on implementation written by Katherine Crowson (but yoinked from the Extra-Samplers repo mentioned above).
* Normal substep merge strategy based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
Thanks!
+1 -1
View File
@@ -113,7 +113,7 @@ class ComposableStepSampler:
"round": False,
},
),
"substeps": ("INT", {"default": 1, "min": 1, "max": 100}),
"substeps": ("INT", {"default": 1, "min": 1, "max": 1000}),
"step_method": (tuple(STEP_SAMPLERS.keys()),),
},
"optional": {
+2 -1
View File
@@ -40,6 +40,7 @@ def composable_sampler(
model_call_cache=None
if "model_call_cache" not in copts
else History(x, copts["model_call_cache"]),
model_call_cache_threshold=copts.get("model_call_cache_threshold", 0),
noise_sampler=noise_sampler,
callback=callback,
eta=eta if eta != 1.0 else copts["eta"],
@@ -60,7 +61,7 @@ def composable_sampler(
samplers += (ssampler,) * sitem["substeps"]
substeps += sitem["substeps"]
msitem = copts["merge_sampler"]
if copts["merge_method"] == "sample":
if copts["merge_method"] in ("sample", "sample_uncached"):
custom_noise = msitem.get("custom_noise_opt")
if custom_noise is None:
curr_ns = noise_sampler
+31 -39
View File
@@ -31,9 +31,7 @@ class NormalMergeSubstepsSampler(MergeSubstepsSampler):
for subidx, ssampler in enumerate(self.samplers):
print(" SUBSTEP", subidx, ssampler)
ss.denoised = ss.model(x, ss.sigma * ss.s_in)
z_k, noise_strength = (
ssampler.step if ss.sigma_next != 0 else ssampler.final_step
)(x, ss)
z_k, noise_strength = ssampler.step(x, ss)
z_avg += renoise_weight * z_k
if ss.sigma_next == 0:
continue
@@ -69,18 +67,17 @@ class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler):
stretch = (ss.sigma - ss.sigma_next) * self.stretch
sig_adj = ss.sigma + stretch
ss = self.ss.clone_edit(sigma=sig_adj)
orig_x = x
x = x + ss.noise_sampler(orig_ss.sigma, ss.sigma_next) * stretch * ss.s_noise
ss.denoised = ss.model(x, sig_adj * ss.s_in)
noise_total = 0.0
for subidx, ssampler in enumerate(self.samplers):
print(" SUBSTEP", subidx, ssampler)
curr_x = x + scale_noise(
curr_x = orig_x + scale_noise(
ssampler.noise_sampler(sig_adj, ss.sigma_next),
ssampler.s_noise * stretch,
)
z_k, noise_strength = (
ssampler.step if ss.sigma_next != 0 else ssampler.final_step
)(curr_x, ss)
z_k, noise_strength = ssampler.step(curr_x, ss)
z_avg += renoise_weight * z_k
if ss.sigma_next == 0:
continue
@@ -99,13 +96,12 @@ class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler):
class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler):
def __init__(
self, ss, samplers, *, merge_sample_skip=False, merge_sampler, **kwargs
):
cache_model = True
def __init__(self, ss, samplers, *, merge_sampler, **kwargs):
super().__init__(ss, samplers, **kwargs)
self.merge_sampler = merge_sampler
self.merge_ss = None
self.merge_sample_skip = merge_sample_skip
def step(self, x):
ss = self.ss
@@ -116,31 +112,25 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler):
ss.denoised = None
stretch = (ss.sigma - ss.sigma_next) * self.stretch
sig_adj = ss.sigma + stretch
if self.merge_sample_skip:
stretch = (ss.sigma - ss.sigma_next) * self.stretch
sig_adj = ss.sigma + stretch
ss.denoised = ss.model(
curr_x,
# + ss.noise_sampler(sig_adj.sigma, ss.sigma_next) * stretch * ss.s_noise,
ss.s_in * sig_adj,
)
ss = self.ss.clone_edit(sigma=sig_adj)
for subidx, ssampler in enumerate(self.samplers):
if not self.merge_sample_skip:
ss.denoised = ss.model(curr_x, ss.s_in * ss.sigma)
else:
curr_x = (
x
+ ssampler.noise_sampler(sig_adj, ss.sigma_next)
* ssampler.s_noise
* stretch
if subidx == 0 or not self.cache_model:
ss.denoised = ss.model(
curr_x,
# + ss.noise_sampler(sig_adj.sigma, ss.sigma_next) * stretch * ss.s_noise,
ss.s_in * sig_adj,
)
print(" SUBSTEP", subidx, ssampler)
z_k, noise_strength = (
ssampler.step if ss.sigma_next != 0 else ssampler.final_step
)(curr_x, ss)
curr_x = (
x
+ ssampler.noise_sampler(sig_adj, ss.sigma_next)
* ssampler.s_noise
* stretch
)
print(" SUBSTEP", subidx, ssampler, stretch)
z_k, noise_strength = ssampler.step(curr_x, ss)
z_avg += renoise_weight * z_k
curr_x = z_k
if not noise_strength or ss.sigma_next == 0 or self.merge_sample_skip:
if not noise_strength or ss.sigma_next == 0:
continue
curr_x += (
ssampler.noise_sampler(ss.sigma, ss.sigma_next)
@@ -165,15 +155,14 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler):
xhist=History(x, 2),
s_noise=msampler.s_noise,
eta=msampler.eta,
model_call_cache=None,
)
else:
merge_ss = self.merge_ss
merge_ss.denoised = result
merge_ss.update(self.ss.idx)
final = merge_ss.sigma_next == 0
merged, noise_strength = (msampler.step if not final else msampler.final_step)(
x, merge_ss
)
merged, noise_strength = msampler.step(x, merge_ss)
if not final:
ss = self.ss
merged = (
@@ -188,6 +177,10 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler):
return merged
class SampleUncachedMergeSubstepsSampler(SampleMergeSubstepsSampler):
cache_model = False
class DivideMergeSubstepsSampler(MergeSubstepsSampler):
def __init__(self, ss, samplers, *, schedule_multiplier=4, **kwargs):
super().__init__(ss, samplers, **kwargs)
@@ -226,9 +219,7 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler):
subss.update(subidx)
subss.denoised = subss.model(x, subss.sigma)
ssampler = samplers[subidx]
x, noise_strength = (
ssampler.step if subss.sigma_next != 0 else ssampler.final_step
)(x, subss)
x, noise_strength = ssampler.step(x, subss)
if not noise_strength or subss.sigma_next == 0:
continue
x = (
@@ -246,7 +237,8 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler):
MERGE_SUBSTEPS_CLASSES = {
"normal": NormalMergeSubstepsSampler,
"divide": DivideMergeSubstepsSampler,
"average": AverageMergeSubstepsSampler,
"sample": SampleMergeSubstepsSampler,
"divide": DivideMergeSubstepsSampler,
"sample_uncached": SampleUncachedMergeSubstepsSampler,
}
+144 -11
View File
@@ -27,7 +27,7 @@ class SingleStepSampler:
raise NotImplementedError
# Euler - based on original ComfyUI implementation
def final_step(self, x, ss):
def euler_step(self, x, ss):
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
d = to_d(x, ss.sigma, ss.denoised)
dt = sigma_down - ss.sigma
@@ -45,12 +45,10 @@ class ReversibleSingleStepSampler(SingleStepSampler):
class EulerStep(SingleStepSampler):
name = "euler"
def step(self, x, ss):
return self.final_step(x, ss)
step = SingleStepSampler.euler_step
class DPMPP2MStep(SingleStepSampler):
class DPMPPStepBase(SingleStepSampler):
@staticmethod
def sigma_fn(t):
return t.neg().exp()
@@ -59,7 +57,11 @@ class DPMPP2MStep(SingleStepSampler):
def t_fn(t):
return t.log().neg()
class DPMPP2MStep(DPMPPStepBase):
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t, t_next = self.t_fn(ss.sigma), self.t_fn(ss.sigma_next)
h = t_next - t
st, st_next = self.sigma_fn(t), self.sigma_fn(t_next)
@@ -80,6 +82,8 @@ class DPMPP2MSDEStep(SingleStepSampler):
self.solver_type = solver_type
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
denoised = ss.denoised
if ss.sigma_next == 0:
return denoised, None
@@ -115,6 +119,8 @@ class DPMPP3MSDEStep(SingleStepSampler):
name = "dpmpp_3m_sde"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
denoised = ss.denoised
if ss.sigma_next == 0:
return denoised, 0
@@ -153,6 +159,8 @@ class ReversibleHeunStep(ReversibleSingleStepSampler):
name = "reversible_heun"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(self.reta)
dt = sigma_down - ss.sigma
@@ -180,6 +188,8 @@ class ReversibleHeun1SStep(ReversibleSingleStepSampler):
name = "reversible_heun_1s"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
# Reversible Heun-inspired update (first-order)
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(self.reta)
@@ -225,6 +235,8 @@ class RESStep(SingleStepSampler):
pass
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
denoised = ss.denoised
lam_next = (
@@ -254,6 +266,8 @@ class TrapezoidalStep(SingleStepSampler):
name = "trapezoidal"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
dt = ss.sigma_next - ss.sigma
denoised = ss.denoised
@@ -282,6 +296,8 @@ class BogackiStep(ReversibleSingleStepSampler):
reversible = False
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
sigma_down_reversible, sigma_up_reversible = ss.get_ancestral_step(self.reta)
sigma, sigma_next = ss.sigma, sigma_down
@@ -329,6 +345,8 @@ class RK4Step(SingleStepSampler):
name = "rk4"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
sigma = ss.sigma
# Calculate the derivative using the model
@@ -388,7 +406,7 @@ class EulerDancingStep(SingleStepSampler):
self.leap = leap
self.dyn_deta_start = dyn_deta_start
self.dyn_deta_end = dyn_deta_end
if dyn_deta_mode not in ("lerp", "deta"):
if dyn_deta_mode not in ("lerp", "lerp_alt", "deta"):
raise ValueError("Bad dyn_deta_mode")
self.dyn_deta_mode = dyn_deta_mode
@@ -407,6 +425,8 @@ class EulerDancingStep(SingleStepSampler):
# Euler method
dt = sigma_down - ss.sigma
x = x + d * dt
if curr_leap == 1:
return x, sigma_up
if None not in (self.dyn_deta_start, self.dyn_deta_end):
if self.dyn_deta_start == self.dyn_deta_end:
dance_scale = self.dyn_deta_start
@@ -421,28 +441,138 @@ class EulerDancingStep(SingleStepSampler):
print("DANCE?", dance_scale, ss.idx, is_danceable, curr_leap)
if not is_danceable or abs(dance_scale) < 1e-04:
return x, sigma_up
orig_x = x
sigma_down_normal, sigma_up_normal = get_ancestral_step(
ss.sigma, ss.sigma_next, self.eta
)
if self.dyn_deta_mode == "lerp":
dt_normal = sigma_down_normal - ss.sigma
x_normal = x + d * dt_normal
else:
x_normal = x
x = x + self.noise_sampler(ss.sigma, sigma_leap) * self.s_noise * sigma_up
sigma_down2, sigma_up2 = get_ancestral_step(
sigma_leap,
ss.sigma_next,
eta=self.deta * (1.0 if self.dyn_deta_mode == "lerp" else dance_scale),
eta=self.deta * (1.0 if self.dyn_deta_mode != "deta" else dance_scale),
)
d_2 = to_d(x, sigma_leap, ss.denoised)
dt_2 = sigma_down2 - sigma_leap
result = x + d_2 * dt_2
noise_diff = sigma_up2 - sigma_up * dance_scale
noise_scale = sigma_up2 + noise_diff * (0.025 * curr_leap)
print(
"DANCE NOISE",
noise_scale,
"--",
sigma_up_normal,
sigma_up_normal - noise_scale,
)
if self.dyn_deta_mode == "deta" or dance_scale == 1.0:
return result, sigma_up2
result = torch.lerp(orig_x, result, dance_scale)
return result, noise_scale
result = torch.lerp(x_normal, result, dance_scale)
# FIXME: Broken for noise samplers that care about s/sn
return result, torch.lerp(sigma_up, sigma_up2, dance_scale)
return result, noise_scale
class DPMPP2SStep(DPMPPStepBase):
name = "dpmpp_2s"
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t_fn, sigma_fn = self.t_fn, self.sigma_fn
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
# DPM-Solver++(2S)
t, t_next = t_fn(ss.sigma), t_fn(sigma_down)
r = 1 / 2
h = t_next - t
s = t + r * h
x_2 = (sigma_fn(s) / sigma_fn(t)) * x - (-h * r).expm1() * ss.denoised
denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=0)
x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_2
return x, sigma_up
class DPMPPSDEStep(DPMPPStepBase):
name = "dpmpp_sde"
def __init__(self, *args, r=1 / 2, **kwargs):
super().__init__(*args, **kwargs)
self.r = r
def step(self, x, ss):
if ss.sigma_next == 0:
return self.euler_step(x, ss)
t_fn, sigma_fn = self.t_fn, self.sigma_fn
r, eta, s_noise = self.r, self.eta, self.s_noise
noise_sampler = self.noise_sampler
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
# DPM-Solver++
t, t_next = t_fn(ss.sigma), t_fn(ss.sigma_next)
h = t_next - t
s = t + h * r
fac = 1 / (2 * r)
# Step 1
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta)
s_ = t_fn(sd)
x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - (t - s_).expm1() * ss.denoised
x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su
denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=0)
# Step 2
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta)
t_next_ = t_fn(sd)
denoised_d = (1 - fac) * ss.denoised + fac * denoised_2
x = (sigma_fn(t_next_) / sigma_fn(t)) * x - (t - t_next_).expm1() * denoised_d
return x, su
# Based on implementation from https://github.com/Clybius/ComfyUI-Extra-Samplers
# Which was originally written by Katherine Crowson
class TTMJVPStep(SingleStepSampler):
name = "ttm_jvp"
def __init__(self, *args, alternate_phi_2_calc=True, **kwargs):
super().__init__(*args, **kwargs)
self.alternate_phi_2_calc = alternate_phi_2_calc
def step(self, x, ss):
if ss.sigma_next == 0:
return ss.denoised, ss.sigma.new_zeros(1)
sigma_down, sigma_up = ss.get_ancestral_step(self.eta)
sigma, sigma_next = ss.sigma, ss.sigma_next
# 2nd order truncated Taylor method
t, s = -sigma.log(), -sigma_next.log()
h = s - t
h_eta = h * (self.eta + 1)
eps = to_d(x, sigma, ss.denoised)
denoised, denoised_prime = ss.model_jvp(
x, sigma, (eps * -sigma, -sigma), model_call_idx=0
)
phi_1 = -torch.expm1(-h_eta)
if self.alternate_phi_2_calc:
phi_2 = torch.expm1(-h) + h # seems to work better with eta > 0
else:
phi_2 = torch.expm1(-h_eta) + h_eta
x = torch.exp(-h_eta) * x + phi_1 * ss.denoised + phi_2 * denoised_prime
if not self.eta:
return x, ss.sigma.new_zeros(1)
phi_1_noise = torch.sqrt(-torch.expm1(-2 * h * self.eta))
return x, sigma_next * phi_1_noise
STEP_SAMPLERS = {
"euler": EulerStep,
"dpmpp_sde": DPMPPSDEStep,
"dpmpp_2m": DPMPP2MStep,
"dpmpp_2m_sde": DPMPP2MSDEStep,
"dpmpp_3m_sde": DPMPP3MSDEStep,
"dpmpp_2s": DPMPP2SStep,
"reversible_heun": ReversibleHeunStep,
"reversible_heun_1s": ReversibleHeun1SStep,
"res": RESStep,
@@ -451,6 +581,7 @@ STEP_SAMPLERS = {
"reversible_bogacki": ReversibleBogackiStep,
"rk4": RK4Step,
"euler_dancing": EulerDancingStep,
"ttm_jvp": TTMJVPStep,
}
__all__ = (
@@ -459,6 +590,7 @@ __all__ = (
"DPMPP2MStep",
"DPMPP2MSDEStep",
"DPMPP3MSDEStep",
"DPMPP2SStep",
"ReversibleHeunStep",
"ReversibleHeun1SStep",
"RESStep",
@@ -466,4 +598,5 @@ __all__ = (
"BogackiStep",
"ReversibleBogackiStep",
"EulerDancingStep",
"TTMJVPStep",
)
+39 -1
View File
@@ -52,6 +52,7 @@ class SamplerState:
callback=None,
denoised=None,
model_call_cache=None,
model_call_cache_threshold=0,
eta=1.0,
reta=1.0,
s_noise=1.0,
@@ -69,6 +70,7 @@ class SamplerState:
self.callback_ = callback
self.noise_sampler = noise_sampler
self.model_call_cache = model_call_cache
self.model_call_cache_threshold = model_call_cache_threshold
self.update(idx)
def update(self, idx=None):
@@ -89,15 +91,50 @@ class SamplerState:
def model(self, x, sigma, *, model_call_idx=-1, **kwargs):
mcc = self.model_call_cache
if mcc is None or model_call_idx < 0 or model_call_idx >= mcc.size:
print("MODEL CALL", model_call_idx)
return self.model_(x, sigma * self.s_in, **self.extra_args, **kwargs)
if model_call_idx < mcc.pos:
if (
model_call_idx < mcc.pos
and model_call_idx >= self.model_call_cache_threshold
):
print("CACHED MODEL CALL", model_call_idx)
return mcc.history[model_call_idx]
result = self.model_(x, sigma * self.s_in, **self.extra_args, **kwargs)
mcc.push(result)
print("CACHING MODEL CALL", model_call_idx, mcc.size, mcc.pos)
return result
def model_jvp(self, x, sigma, tangents, *, model_call_idx=-1, **kwargs):
def call_model(x, sigma):
return self.model_(x, sigma * self.s_in, **self.extra_args)
mcc = self.model_call_cache
if mcc is not None:
if not hasattr(self, "model_jvp_call_cache"):
self.model_jvp_call_cache = History(x, mcc.size)
jmcc = self.model_jvp_call_cache
if mcc is None or model_call_idx < 0 or model_call_idx >= mcc.size:
print("JMODEL CALL", model_call_idx)
return torch.func.jvp(call_model, (x, sigma), tangents)
if mcc.size != jmcc.size or mcc.pos != jmcc.pos:
raise RuntimeError(
"TTM JVP steps with caching enabled currently must run first"
)
if (
model_call_idx < mcc.pos
and model_call_idx >= self.model_call_cache_threshold
):
print("JCACHED MODEL CALL", model_call_idx)
return mcc.history[model_call_idx], jmcc.history[model_call_idx]
denoised, denoised_prime = torch.func.jvp(call_model, (x, sigma), tangents)
mcc.push(denoised)
jmcc.push(denoised_prime)
print("JCACHING MODEL CALL", model_call_idx, mcc.size, mcc.pos)
return denoised, denoised_prime
def get_ancestral_step(self, eta=1.0):
return get_ancestral_step(self.sigma, self.sigma_next, eta=eta)
@@ -108,6 +145,7 @@ class SamplerState:
"dhist",
"xhist",
"model_call_cache",
"model_call_cache_threshold",
"extra_args",
"s_in",
"eta",