From be5e5d8933aa7f48182a3d81a80c00e6a8d68bbb Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 30 May 2024 19:47:44 -0600 Subject: [PATCH] All sorts of fun stuff --- README.md | 28 ++++++-- py/nodes.py | 2 +- py/sampling.py | 3 +- py/substep_merging.py | 70 +++++++++---------- py/substep_samplers.py | 155 ++++++++++++++++++++++++++++++++++++++--- py/substep_sampling.py | 40 ++++++++++- 6 files changed, 238 insertions(+), 60 deletions(-) diff --git a/README.md b/README.md index 7927e92..4ca31a2 100644 --- a/README.md +++ b/README.md @@ -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! diff --git a/py/nodes.py b/py/nodes.py index dbbea5a..3b5b764 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -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": { diff --git a/py/sampling.py b/py/sampling.py index da58967..bfea6ed 100644 --- a/py/sampling.py +++ b/py/sampling.py @@ -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 diff --git a/py/substep_merging.py b/py/substep_merging.py index c26aefa..5a2d1f4 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -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, } diff --git a/py/substep_samplers.py b/py/substep_samplers.py index 491e9ee..303dd80 100644 --- a/py/substep_samplers.py +++ b/py/substep_samplers.py @@ -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", ) diff --git a/py/substep_sampling.py b/py/substep_sampling.py index 0cd6a6c..2ebfccc 100644 --- a/py/substep_sampling.py +++ b/py/substep_sampling.py @@ -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",