Add lookahead substep merge method

Rename normal substep merge method to supreme_avg
This commit is contained in:
blepping
2024-09-03 17:32:12 -06:00
parent cc30fde46c
commit d65c096438
3 changed files with 106 additions and 10 deletions
+18 -1
View File
@@ -229,8 +229,9 @@ When running multiple substeps per step, the results will combined based on the
* `simple`: Doesn't merge anything: only runs a single substep per step.
* `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 step (and possibly additional times for higher order samplers). Each substep shares the first model call result. The results are averaged together. *Note*: Since the first model call is shared and the initial input is the same for each substep, there is no point in running multiple identical substeps. Also note: This merge strategy doesn't work well with non-ancestral samplers (i.e. dpmpp_2m or any sampler with `eta: 0`).
* `supreme_avg`: The model is called at least once per step (and possibly additional times for higher order samplers). Each substep shares the first model call result. The results are averaged together. *Note*: Since the first model call is shared and the initial input is the same for each substep, there is no point in running multiple identical substeps. Also note: This merge strategy doesn't work well with non-ancestral samplers (i.e. dpmpp_2m or any sampler with `eta: 0`).
* `overshoot`: The model is called at least once per step. It will sample steps equal to the number of substeps, starting from the current step. Then it will restart back to the expected step.
* `lookahead`: Similar to `overshoot`, it samples ahead based on the number of substeps. The last model prediction is used to do a Euler step to the expected step. *Note:* Very experimental, likely to change in the future.
<!--
* `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 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.
@@ -260,6 +261,7 @@ The left side group matches steps 0, 1, 2. The right side group matches all step
-->
* `restart_custom_noise`: Currently only used by the `overshoot` merge method.
* `custom_noise`: Currently used by the `lookahead` merge method.
#### Text Parameters
@@ -302,6 +304,21 @@ restart:
immiscible:
size: 0
# Only used by the lookahead merge method currently.
lookahead:
# Works like normal samplers, essentially. Disabled by default.
eta: 0.0
# Scales the noise added by lookahead sampling.
s_noise: 1.0
# Controls how much noise to remove in the prediction phase. Higher values will remove more noise.
dt_factor: 1.0
# Immiscible block same as described above.
immiscible:
size: 0
pre_filter: null
post_filter: null
+88 -6
View File
@@ -7,6 +7,7 @@ from . import expression as expr
from . import utils
from .filtering import make_filter, FilterRefs
from .noise import ImmiscibleNoise
from .restart import Restart
from .step_samplers import STEP_SAMPLERS
from .utils import check_time, fallback
@@ -136,8 +137,8 @@ class SimpleSubstepsSampler(MergeSubstepsSampler):
return self.merge_steps(sr.noise_x(ss=ss))
class NormalMergeSubstepsSampler(MergeSubstepsSampler):
name = "normal"
class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
name = "supreme_avg"
def step(self, x):
ss = self.ss
@@ -514,13 +515,94 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
return x
class LookaheadMergeSubstepsSampler(MergeSubstepsSampler):
name = "lookahead"
def __init__(self, ss, group, **kwargs):
super().__init__(ss, group, **kwargs)
lookahead = self.options.pop("lookahead", {})
self.lookahead_eta = lookahead.pop("eta", 0.0)
self.lookahead_s_noise = lookahead.pop("s_noise", 1.0)
self.lookahead_dt_factor = lookahead.pop("dt_factor", 1.0)
immiscible = lookahead.get("immiscible", False)
self.immiscible = (
ImmiscibleNoise(**immiscible) if immiscible is not False else False
)
self.custom_noise = self.options.get("custom_noise")
def step(self, x):
orig_x = x.clone()
ss = self.ss
subss = self.ss.clone_edit(idx=ss.idx, sigmas=ss.sigmas)
substep = 0
max_idx = len(ss.sigmas) - 1
eff_substeps = min(max_idx - ss.idx, self.substeps)
pbar = tqdm.tqdm(total=eff_substeps, initial=0, disable=ss.disable_status)
for ssampler in self.samplers:
substeps_remain = eff_substeps - substep
if substeps_remain == 0:
break
custom_noise = ssampler.options.get(
"custom_noise", self.options.get("custom_noise")
)
noise_sampler = ss.noise.make_caching_noise_sampler(
custom_noise,
ssampler.max_noise_samples(),
ss.sigma,
ss.sigma_next,
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
)
ssampler.noise_sampler = noise_sampler
for subidx in range(min(substeps_remain, ssampler.substeps)):
subss.update(ss.idx + substep, substep=substep)
pbar.set_description(
f"substep({ssampler.name}): {subss.sigma.item():.03} -> {subss.sigma_next.item():.03}"
)
subss.hist.push(subss.model(x, subss.sigma, ss=subss))
subss.refs = FilterRefs.from_ss(subss, have_current=True)
if substep == 0:
self.callback(ss=subss)
sr = self.simple_substep(x, ssampler, ss=subss)
x = sr.x
noise_strength = sr.noise_scale
if noise_strength != 0 and subss.sigma_next != 0:
x = sr.noise_x(ss=subss)
substep += 1
pbar.update(1)
if substeps_remain == 1:
break
pbar.update(0)
sigma_down, sigma_up = ss.get_ancestral_step(
eta=self.lookahead_eta, sigma=ss.sigma, sigma_next=ss.sigma_next
)
if sr.sigma_next == sigma_down:
return x
dt = (
torch.sqrt(1.0 + (ss.sigma - sigma_down) ** 2) * 0.05
+ (ss.sigma - sigma_down) * 0.95
) * self.lookahead_dt_factor
denoised = sr.denoised
d = (orig_x - denoised) / ss.sigma
x = orig_x + d * -dt
if sigma_down == 0 or sigma_up == 0:
return x
noise_sampler = ss.noise.make_caching_noise_sampler(
self.custom_noise,
1,
ss.sigma,
ss.sigma_next,
immiscible=fallback(self.immiscible, ss.noise.immiscible),
)
x += ss.noise.scale_noise(noise_sampler(refs=ss.refs), sigma_up)
return x
MERGE_SUBSTEPS_CLASSES = {
"default (simple)": SimpleSubstepsSampler,
"normal": NormalMergeSubstepsSampler,
"supreme_avg": SupremeAvgMergeSubstepsSampler,
"divide": DivideMergeSubstepsSampler,
"overshoot": OvershootMergeSubstepsSampler,
# "average": AverageMergeSubstepsSampler,
# "sample": SampleMergeSubstepsSampler,
# "sample_uncached": SampleUncachedMergeSubstepsSampler,
"simple": SimpleSubstepsSampler,
"lookahead": LookaheadMergeSubstepsSampler,
}
-3
View File
@@ -4,9 +4,6 @@ import torch
from comfy.k_diffusion.sampling import to_d
from . import latent
# def scale_noise_(
# noise,
# factor=1.0,