diff --git a/README.md b/README.md index 4ba120c..26130dd 100644 --- a/README.md +++ b/README.md @@ -16,8 +16,11 @@ 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. 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. +* `model_call_cache`(unset): Caches the result of model calls. For example, Bogacki is 3 model calls per step. The first one usually depends on the merge strategy: `average` for example shares the first model evaluation between substeps, but subsequent model calls (i.e. Bogacki 2nd and 3rd model evaluations) still occur. When the model call cache is active, it's possible to cache those evalutions and avoid a model call for the remaining substeps. 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`(`1`): Disables caching model call results with a call index below the threshold value (starting at 0). For example, if set to `2` and using a sampler like Bogacki that calls the model two extra times, the first will never be cached. The default value of `1` disabling caching for the first model call per substep. I generally would not recommend setting it to `0`, especially with `average` or `sample` merge strategies. +* `model_call_cache_max_use`(`1000000`): The number of times cache items can be re-used. The default is effectively no limit. Where would this be useful? Let's say you're using the `average` merge strategy and a multi step sampler that calls the model at least one more time with 50 substeps. If you set the value to `25`, the model cache result will be updated around substep 25 which _may_ produce better results than reusing the result 50 times. + +Since it's kind of confusing even for me, a little more explanation: The model call cache caches results for model call indexes between `model_call_cache_threshold` and `model_call_cache - 1`. If you set `model_call_cache_threshold` to `0` and `model_call_cache` to `1` then only the first model call will be cached. If you set `model_call_cache_threshold` to `1` and `model_call_cache` to `2` then call 0 will not be cached, call 1 will be cached, call 2 will be cached, call 3 will not be cached, and so on. #### Merging diff --git a/py/sampling.py b/py/sampling.py index bfea6ed..d373c20 100644 --- a/py/sampling.py +++ b/py/sampling.py @@ -3,7 +3,7 @@ from tqdm.auto import trange from .substep_samplers import STEP_SAMPLERS -from .substep_sampling import SamplerState, History +from .substep_sampling import SamplerState, History, ModelCallCache from .substep_merging import MERGE_SUBSTEPS_CLASSES @@ -29,24 +29,6 @@ def composable_sampler( def noise_sampler(_s, _sn): return torch.randn_like(x) - ss = SamplerState( - model, - sigmas, - 0, - x.new_ones((x.shape[0],)), - History(x, 3), - History(x, 2), - extra_args, - 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"], - s_noise=s_noise if s_noise != 1.0 else copts["s_noise"], - reta=copts.get("reta", 1.0), - ) samplers = [] substeps = 0 for sitem in copts["chain"].items: @@ -58,8 +40,9 @@ def composable_sampler( x, sigmas[-1], sigmas[0], normalized=True ) ssampler = STEP_SAMPLERS[sitem["step_method"]](noise_sampler=curr_ns, **sitem) - samplers += (ssampler,) * sitem["substeps"] - substeps += sitem["substeps"] + samplers.append(ssampler) + # samplers += (ssampler,) * sitem["substeps"] + substeps += ssampler.substeps msitem = copts["merge_sampler"] if copts["merge_method"] in ("sample", "sample_uncached"): custom_noise = msitem.get("custom_noise_opt") @@ -75,6 +58,27 @@ def composable_sampler( pass else: merge_sampler = None + ss = SamplerState( + ModelCallCache( + model, + x, + x.new_ones((x.shape[0],)), + extra_args, + size=copts.get("model_call_cache", 0), + max_use=copts.get("model_call_cache_max_use", 1000000), + threshold=copts.get("model_call_cache_threshold", 0), + ), + sigmas, + 0, + History(x, 3), + History(x, 2), + extra_args, + noise_sampler=noise_sampler, + callback=callback, + eta=eta if eta != 1.0 else copts["eta"], + s_noise=s_noise if s_noise != 1.0 else copts["s_noise"], + reta=copts.get("reta", 1.0), + ) merge_sampler = MERGE_SUBSTEPS_CLASSES[copts["merge_method"]]( ss, samplers, @@ -83,7 +87,6 @@ def composable_sampler( for idx in trange(len(sigmas) - 1, disable=disable): print(f"STEP {idx+1}") ss.update(idx) - if ss.model_call_cache is not None: - ss.model_call_cache.reset() + ss.model.reset_cache() x = merge_sampler.step(x) return x diff --git a/py/substep_merging.py b/py/substep_merging.py index e381598..6f98a56 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -8,6 +8,7 @@ class MergeSubstepsSampler: def __init__(self, ss, samplers, **_kwargs): self.ss = ss self.samplers = samplers + self.substeps = sum(sampler.substeps for sampler in samplers) def step(self, x): raise NotImplementedError @@ -23,22 +24,24 @@ class NormalMergeSubstepsSampler(MergeSubstepsSampler): def step(self, x): ss = self.ss - substeps = len(self.samplers) + substeps = self.substeps renoise_weight = 1.0 / substeps z_avg = torch.zeros_like(x) noise = torch.zeros_like(x) noise_total = 0.0 - for subidx, ssampler in enumerate(self.samplers): - print(" SUBSTEP", subidx, ssampler.name) - ss.denoised = ss.model(x, ss.sigma * ss.s_in) + for idx, ssampler in enumerate( + sampler for sampler in self.samplers for _ in range(sampler.substeps) + ): + print(f" SUBSTEP {idx+1}: {ssampler.name}") + ss.denoised = ss.model(x, ss.sigma) z_k, noise_strength = ssampler.step(x, ss) z_avg += renoise_weight * z_k - if ss.sigma_next == 0: - continue noise_strength *= ssampler.s_noise + if ss.sigma_next == 0 or noise_strength == 0: + continue noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next) x = z_k - if subidx != substeps - 1: + if idx != substeps - 1: x += noise_curr * noise_strength noise_total += noise_strength.item() * renoise_weight noise += noise_curr * noise_strength @@ -60,7 +63,7 @@ class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler): def step(self, x): ss = orig_ss = self.ss - substeps = len(self.samplers) + substeps = self.substeps renoise_weight = 1.0 / substeps z_avg = torch.zeros_like(x) noise = torch.zeros_like(x) @@ -69,22 +72,26 @@ class AverageMergeSubstepsSampler(NormalMergeSubstepsSampler): 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) + ss.denoised = ss.model(x, sig_adj) noise_total = 0.0 - for subidx, ssampler in enumerate(self.samplers): - print(" SUBSTEP", subidx, ssampler.name) - curr_x = orig_x + scale_noise( - ssampler.noise_sampler(sig_adj, ss.sigma_next), - ssampler.s_noise * stretch, + for idx, ssampler in enumerate(self.samplers): + print( + f" SUBSTEP {idx+1} .. {idx+ssampler.substeps}: {ssampler.name}, stretch={stretch}" ) - z_k, noise_strength = ssampler.step(curr_x, ss) - z_avg += renoise_weight * z_k - if ss.sigma_next == 0: - continue - noise_strength *= ssampler.s_noise - noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next) - noise_total += noise_strength.item() * renoise_weight - noise += noise_curr * noise_strength + for sidx in range(ssampler.substeps): + curr_x = orig_x + scale_noise( + ssampler.noise_sampler(sig_adj, ss.sigma_next), stretch + ) + z_k, noise_strength = ssampler.step(curr_x, ss) + z_avg += renoise_weight * z_k + if ss.sigma_next == 0: + continue + noise_strength *= ssampler.s_noise + if noise_strength == 0: + continue + noise_curr = ssampler.noise_sampler(ss.sigma, ss.sigma_next) + noise_total += noise_strength.item() * renoise_weight + noise += noise_curr * noise_strength ss.dhist.push(ss.denoised) ss.denoised = None x = self.merge_steps(x, z_avg) @@ -105,7 +112,7 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): def step(self, x): ss = self.ss - substeps = len(self.samplers) + substeps = self.substeps renoise_weight = 1.0 / substeps z_avg = torch.zeros_like(x) curr_x = x @@ -113,30 +120,33 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): stretch = (ss.sigma - ss.sigma_next) * self.stretch sig_adj = ss.sigma + stretch ss = self.ss.clone_edit(sigma=sig_adj) - for subidx, ssampler in enumerate(self.samplers): - 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, + for idx, ssampler in enumerate(self.samplers): + print( + f" SUBSTEP {idx+1} .. {idx+ssampler.substeps}: {ssampler.name}, stretch={stretch}" + ) + for sidx in range(ssampler.substeps): + if idx == 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, + sig_adj, + ) + curr_x = ( + x + + ssampler.noise_sampler(sig_adj, ss.sigma_next) + * ssampler.s_noise + * stretch + ) + z_k, noise_strength = ssampler.step(curr_x, ss) + z_avg += renoise_weight * z_k + curr_x = z_k + if noise_strength == 0 or ss.sigma_next == 0: + continue + curr_x += ( + ssampler.noise_sampler(ss.sigma, ss.sigma_next) + * ssampler.s_noise + * noise_strength ) - curr_x = ( - x - + ssampler.noise_sampler(sig_adj, ss.sigma_next) - * ssampler.s_noise - * stretch - ) - print(" SUBSTEP", subidx, ssampler.name, 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: - continue - curr_x += ( - ssampler.noise_sampler(ss.sigma, ss.sigma_next) - * ssampler.s_noise - * noise_strength - ) ss.dhist.push(ss.denoised) ss.denoised = None x = self.merge_steps(curr_x, z_avg) @@ -145,8 +155,7 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): return x def merge_steps(self, x, result): - if self.ss.model_call_cache is not None: - self.ss.model_call_cache.reset() + self.ss.model.reset_cache() msampler = self.merge_sampler if self.merge_ss is None: merge_ss = self.merge_ss = self.ss.clone_edit( @@ -155,7 +164,7 @@ class SampleMergeSubstepsSampler(AverageMergeSubstepsSampler): xhist=History(x, 2), s_noise=msampler.s_noise, eta=msampler.eta, - model_call_cache=None, + # model_call_cache=None, ) else: merge_ss = self.merge_ss @@ -186,10 +195,7 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler): super().__init__(ss, samplers, **kwargs) self.schedule_multiplier = schedule_multiplier - def step(self, x): - ss = self.ss - samplers = self.samplers - substeps = len(samplers) + def make_schedule(self, ss): max_steps = len(self.ss.sigmas) - 1 sigmas_slice = ss.sigmas[ ss.idx : min(max_steps + 1, ss.idx + self.schedule_multiplier) @@ -203,24 +209,30 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler): torch.linspace( sigmas_slice[idx], sigmas_slice[idx + 1], - steps=substeps + 1, + steps=self.substeps + 1, device=sigmas_slice.device, dtype=sigmas_slice.dtype, )[0 if not idx else 1 :] for idx in range(len(sigmas_slice) - 1) ) # print("CHUNKS", chunks) - subsigmas = torch.cat(chunks) + return torch.cat(chunks) + + def step(self, x): + ss = self.ss # print("SUBSIGMAS", subsigmas) - subss = self.ss.clone_edit(idx=0, sigmas=subsigmas) + subss = self.ss.clone_edit(idx=0, sigmas=self.make_schedule(ss)) subss.main_idx = ss.idx subss.main_sigmas = ss.sigmas - for subidx in range(substeps): - subss.update(subidx) + + for idx, ssampler in enumerate( + sampler for sampler in self.samplers for _ in range(sampler.substeps) + ): + print(f" SUBSTEP {idx+1}: {ssampler.name}") + subss.update(idx) subss.denoised = subss.model(x, subss.sigma) - ssampler = samplers[subidx] x, noise_strength = ssampler.step(x, subss) - if not noise_strength or subss.sigma_next == 0: + if noise_strength == 0 or subss.sigma_next == 0: continue x = ( x diff --git a/py/substep_samplers.py b/py/substep_samplers.py index 1163ff7..a5a13dd 100644 --- a/py/substep_samplers.py +++ b/py/substep_samplers.py @@ -18,6 +18,7 @@ class SingleStepSampler: self, *, noise_sampler=None, + substeps=1, s_noise=1.0, eta=1.0, dyn_eta_start=None, @@ -31,6 +32,7 @@ class SingleStepSampler: self.dyn_eta_end = dyn_eta_end self.noise_sampler = noise_sampler self.weight = weight + self.substeps = substeps self.kwargs = kwargs def step(self, x, ss): @@ -207,7 +209,7 @@ class ReversibleHeunStep(ReversibleSingleStepSampler): x_pred = x + d * dt # Denoised sample at the next sigma - denoised_next = ss.model(x_pred, sigma_down, model_call_idx=0) + denoised_next = ss.model(x_pred, sigma_down, model_call_idx=1) # Calculate the derivative at the next sigma d_next = to_d(x_pred, sigma_down, denoised_next) @@ -241,7 +243,7 @@ class ReversibleHeun1SStep(ReversibleSingleStepSampler): sigma_i, ss.dhist[-1] if len(ss.dhist) - else ss.model(eff_x, sigma_i, model_call_idx=0), + else ss.model(eff_x, sigma_i, model_call_idx=1), ) # Predict the sample at the next sigma using Euler step @@ -289,7 +291,7 @@ class RESStep(SingleStepSampler): lam_2 = lam + c2_h sigma_2 = lam_2.neg().exp() - denoised2 = ss.model(x_2, sigma_2, model_call_idx=0) + denoised2 = ss.model(x_2, sigma_2, model_call_idx=1) x = math.exp(-h) * x + h * (b1 * denoised + b2 * denoised2) return x, sigma_up @@ -313,7 +315,7 @@ class TrapezoidalStep(SingleStepSampler): x_pred = x + d_i * dt # Denoised sample at the next sigma - denoised_next = ss.model(x_pred, ss.sigma_next, model_call_idx=0) + denoised_next = ss.model(x_pred, ss.sigma_next, model_call_idx=1) # Calculate the derivative at the next sigma d_next = to_d(x_pred, ss.sigma_next, denoised_next) @@ -350,7 +352,7 @@ class BogackiStep(ReversibleSingleStepSampler): to_d( x + k1 / 2, sigma + dt / 2, - ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=0), + ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=1), ) * dt ) @@ -358,7 +360,7 @@ class BogackiStep(ReversibleSingleStepSampler): to_d( x + 3 * k1 / 4 + k2 / 4, sigma + 3 * dt / 4, - ss.model(x + 3 * k1 / 4 + k2 / 4, sigma + 3 * dt / 4, model_call_idx=1), + ss.model(x + 3 * k1 / 4 + k2 / 4, sigma + 3 * dt / 4, model_call_idx=2), ) * dt ) @@ -395,7 +397,7 @@ class RK4Step(SingleStepSampler): to_d( x + k1 / 2, sigma + dt / 2, - ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=0), + ss.model(x + k1 / 2, sigma + dt / 2, model_call_idx=1), ) * dt ) @@ -403,7 +405,7 @@ class RK4Step(SingleStepSampler): to_d( x + k2 / 2, sigma + dt / 2, - ss.model(x + k2 / 2, sigma + dt / 2, model_call_idx=1), + ss.model(x + k2 / 2, sigma + dt / 2, model_call_idx=2), ) * dt ) @@ -411,7 +413,7 @@ class RK4Step(SingleStepSampler): to_d( x + k3, sigma + dt, - ss.model(x + k3, sigma + dt, model_call_idx=2), + ss.model(x + k3, sigma + dt, model_call_idx=3), ) * dt ) @@ -537,7 +539,7 @@ class DPMPPSDEStep(DPMPPStepBase): 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) + denoised_2 = ss.model(x_2, sigma_fn(s), model_call_idx=1) # Step 2 sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta) @@ -568,8 +570,8 @@ class TTMJVPStep(SingleStepSampler): h_eta = h * (eta + 1) eps = to_d(x, sigma, ss.denoised) - denoised, denoised_prime = ss.model_jvp( - x, sigma, (eps * -sigma, -sigma), model_call_idx=0 + denoised, denoised_prime = ss.model( + x, sigma, tangents=(eps * -sigma, -sigma), model_call_idx=1 ) phi_1 = -torch.expm1(-h_eta) diff --git a/py/substep_sampling.py b/py/substep_sampling.py index 2ebfccc..ec2836a 100644 --- a/py/substep_sampling.py +++ b/py/substep_sampling.py @@ -37,13 +37,75 @@ class History: self.last = None +class ModelCallCache: + def __init__( + self, model, x, s_in, extra_args, *, size=0, max_use=1000000, threshold=1 + ): + self.size = size + self.model = model + self.threshold = threshold + self.s_in = s_in + self.extra_args = extra_args + self.max_use = max_use + if self.size < 1: + return + self.mcc = torch.zeros(size, *x.shape, device=x.device, dtype=x.dtype) + self.jmcc = torch.zeros_like(self.mcc) + self.reset_cache() + + def reset_cache(self): + size = self.size + self.slot = [None] * size + self.jslot = [None] * size + self.slot_use = [self.max_use] * size + + def get(self, idx, *, jvp=False): + idx -= self.threshold + if ( + idx >= self.size + or idx < 0 + or self.slot[idx] is None + or self.slot_use[idx] < 1 + ): + return None + if jvp and self.jslot[idx] is None: + return None + self.slot_use[idx] -= 1 + return self.slot[idx] if not jvp else (self.slot[idx], self.jslot[idx]) + + def set(self, idx, denoised, jdenoised=None): + idx -= self.threshold + if idx < 0 or idx >= self.size: + return + self.slot_use[idx] = self.max_use + self.slot[idx] = denoised + self.jslot[idx] = jdenoised + + def call_model(self, x, sigma, **kwargs): + return self.model(x, sigma * self.s_in, **self.extra_args, **kwargs) + + def __call__(self, x, sigma, *, model_call_idx=0, tangents=None, **kwargs): + result = self.get(model_call_idx, jvp=tangents is not None) + # print( + # f"MODEL: idx={model_call_idx}, size={self.size}, threshold={self.threshold}, cached={result is not None}" + # ) + if result is not None: + return result + if tangents is None: + denoised = self.call_model(x, sigma, **kwargs) + self.set(model_call_idx, denoised) + return denoised + denoised, denoised_prime = torch.func.jvp(self.call_model, (x, sigma), tangents) + self.set(model_call_idx, denoised, jdenoised=denoised_prime) + return denoised, denoised_prime + + class SamplerState: def __init__( self, model, sigmas, idx, - s_in, dhist, xhist, extra_args, @@ -51,17 +113,14 @@ class SamplerState: noise_sampler, callback=None, denoised=None, - model_call_cache=None, - model_call_cache_threshold=0, eta=1.0, reta=1.0, s_noise=1.0, ): - self.model_ = model + self.model = model self.dhist = dhist self.xhist = xhist self.extra_args = extra_args - self.s_in = s_in self.eta = eta self.reta = reta self.s_noise = s_noise @@ -69,8 +128,6 @@ class SamplerState: self.denoised = denoised 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): @@ -88,66 +145,16 @@ class SamplerState: self.sigma, self.sigma_next, eta=self.reta ) - 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 - 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) def clone_edit(self, **kwargs): obj = self.__class__.__new__(self.__class__) for k in ( - "model_", + "model", "dhist", "xhist", - "model_call_cache", - "model_call_cache_threshold", "extra_args", - "s_in", "eta", "reta", "s_noise",