Rewrite model call caching

This commit is contained in:
blepping
2024-05-31 13:21:13 -06:00
parent bf8d525db2
commit b5a4bf9895
5 changed files with 182 additions and 155 deletions
+5 -2
View File
@@ -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
+26 -23
View File
@@ -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
+72 -60
View File
@@ -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
+14 -12
View File
@@ -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)
+65 -58
View File
@@ -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",