Allow overriding CFG scale in parameters
Try to make sure uncond gets generated when needed
This commit is contained in:
@@ -93,6 +93,8 @@ class FilterRefs:
|
||||
"step_pct": float(ss.step / ss.total_steps),
|
||||
"total_steps": ss.total_steps,
|
||||
"sampling_pct": (999 - ms.timestep(ss.sigma).item()) / 999,
|
||||
"is_rectified_flow": ss.model.is_rectified_flow,
|
||||
"original_cfg_scale": ss.model.inner_cfg_scale,
|
||||
})
|
||||
if have_current and len(ss.hist) > 0:
|
||||
fr |= cls.from_mr(ss.hcur)
|
||||
|
||||
+60
-29
@@ -104,14 +104,15 @@ class ModelCallCache:
|
||||
def __init__(
|
||||
self,
|
||||
model,
|
||||
x,
|
||||
s_in,
|
||||
extra_args,
|
||||
x: torch.Tensor,
|
||||
s_in: torch.Tensor,
|
||||
extra_args: dict,
|
||||
*,
|
||||
cache=None,
|
||||
filter=None,
|
||||
cfg1_uncond_optimization=False,
|
||||
):
|
||||
cache: None | dict = None,
|
||||
filter: None | dict = None,
|
||||
cfg1_uncond_optimization: bool = False,
|
||||
cfg_scale_override: None | int | float = None,
|
||||
) -> None:
|
||||
self.cache = ModelCallCacheConfig(**fallback(cache, {}))
|
||||
filtargs = fallback(filter, {}).copy()
|
||||
self.filters = {}
|
||||
@@ -124,6 +125,7 @@ class ModelCallCache:
|
||||
self.s_in = s_in
|
||||
self.extra_args = extra_args
|
||||
self.cfg1_uncond_optimization = cfg1_uncond_optimization
|
||||
self.cfg_scale_override = cfg_scale_override
|
||||
self.is_rectified_flow = x.shape[1] == 16 and isinstance(
|
||||
model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST
|
||||
)
|
||||
@@ -131,13 +133,17 @@ class ModelCallCache:
|
||||
return
|
||||
self.reset_cache()
|
||||
|
||||
def maybe_filter(self, name, latent, *args, **kwargs):
|
||||
def maybe_filter(
|
||||
self, name: str, latent: torch.Tensor, *args: list, **kwargs: dict
|
||||
) -> torch.Tensor:
|
||||
filt = self.filters.get(name)
|
||||
if filt is None:
|
||||
return latent
|
||||
return filt.apply(latent, *args, **kwargs)
|
||||
|
||||
def filter_result(self, result, *args, **kwargs):
|
||||
def filter_result(
|
||||
self, result: ModelResult, *args: list, **kwargs: dict
|
||||
) -> ModelResult:
|
||||
if not self.filters:
|
||||
return result
|
||||
result = result.clone()
|
||||
@@ -153,17 +159,17 @@ class ModelCallCache:
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _fr_add_mr(fr, mr):
|
||||
def _fr_add_mr(fr: filtering.FilterRefs, mr: ModelResult) -> filtering.FilterRefs:
|
||||
frmr = filtering.FilterRefs.from_mr(mr)
|
||||
fr.kvs |= {f"{k}_curr": v for k, v in frmr.kvs.items()}
|
||||
return fr
|
||||
|
||||
def reset_cache(self):
|
||||
def reset_cache(self) -> None:
|
||||
size = self.cache.size
|
||||
self.slot = [None] * size
|
||||
self.slot_use = [self.cache.max_use] * size
|
||||
|
||||
def get(self, idx, *, jvp=False):
|
||||
def get(self, idx: int, *, jvp: bool = False) -> None | ModelResult:
|
||||
idx -= self.cache.threshold
|
||||
if (
|
||||
idx >= self.cache.size
|
||||
@@ -178,32 +184,52 @@ class ModelCallCache:
|
||||
self.slot_use[idx] -= 1
|
||||
return result
|
||||
|
||||
def set(self, idx, mr):
|
||||
def set(self, idx: int, mr: ModelResult) -> None:
|
||||
idx -= self.cache.threshold
|
||||
if idx < 0 or idx >= self.cache.size:
|
||||
return
|
||||
self.slot_use[idx] = self.cache.max_use
|
||||
self.slot[idx] = mr
|
||||
|
||||
def call_model(self, x, sigma, **kwargs):
|
||||
def call_model(
|
||||
self, x: torch.Tensor, sigma: torch.Tensor, **kwargs: dict
|
||||
) -> torch.Tensor:
|
||||
return self.model(x, sigma * self.s_in, **self.extra_args | kwargs)
|
||||
|
||||
@property
|
||||
def model_sampling(self):
|
||||
return self.model.inner_model.inner_model.model_sampling
|
||||
|
||||
@property
|
||||
def inner_cfg_scale(self) -> None | int | float:
|
||||
maybe_cfg_scale = getattr(self.model.inner_model, "cfg", None)
|
||||
return maybe_cfg_scale if isinstance(maybe_cfg_scale, (int, float)) else None
|
||||
|
||||
def set_inner_cfg_scale(self, scale: None | int | float) -> None | int | float:
|
||||
eff_scale = self.cfg_scale_override
|
||||
if scale is not None:
|
||||
eff_scale = None if scale < 0 else scale
|
||||
if eff_scale is None or eff_scale < 0:
|
||||
return None
|
||||
curr_cfg_scale = self.inner_cfg_scale
|
||||
if curr_cfg_scale is None:
|
||||
return None
|
||||
self.model.inner_model.cfg = eff_scale
|
||||
return curr_cfg_scale
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
x,
|
||||
sigma,
|
||||
x: torch.Tensor,
|
||||
sigma: torch.Tensor,
|
||||
*,
|
||||
call_index=0,
|
||||
call_index: int = 0,
|
||||
ss,
|
||||
s_in=None,
|
||||
tangents=None,
|
||||
return_cached=False,
|
||||
require_uncond: bool = False,
|
||||
cfg_scale_override: None | int = None,
|
||||
**kwargs,
|
||||
):
|
||||
) -> ModelResult:
|
||||
filter_refs = ss.refs | filtering.FilterRefs({
|
||||
"model_call": call_index,
|
||||
"orig_x": x,
|
||||
@@ -215,7 +241,7 @@ class ModelCallCache:
|
||||
if result is not None:
|
||||
self._fr_add_mr(filter_refs, result)
|
||||
result = self.filter_result(result, default_ref=x, refs=filter_refs)
|
||||
return (result, True) if return_cached else result
|
||||
return result
|
||||
|
||||
comfy.model_management.throw_exception_if_processing_interrupted()
|
||||
|
||||
@@ -230,13 +256,16 @@ class ModelCallCache:
|
||||
denoised_uncond = denoised_cond
|
||||
return args["denoised"]
|
||||
|
||||
extra_args = self.extra_args | {
|
||||
"model_options": comfy.model_patcher.set_model_options_post_cfg_function(
|
||||
model_options,
|
||||
postcfg,
|
||||
disable_cfg1_optimization=not self.cfg1_uncond_optimization,
|
||||
)
|
||||
}
|
||||
orig_cfg_scale = self.set_inner_cfg_scale(cfg_scale_override)
|
||||
|
||||
model_options = comfy.model_patcher.set_model_options_post_cfg_function(
|
||||
model_options,
|
||||
postcfg,
|
||||
disable_cfg1_optimization=require_uncond
|
||||
or not self.cfg1_uncond_optimization,
|
||||
)
|
||||
|
||||
extra_args = self.extra_args | {"model_options": model_options}
|
||||
s_in = fallback(s_in, self.s_in)
|
||||
x = self.maybe_filter("input", x, refs=filter_refs)
|
||||
|
||||
@@ -245,6 +274,7 @@ class ModelCallCache:
|
||||
|
||||
if tangents is None:
|
||||
denoised = call_model(x, sigma, **kwargs)
|
||||
self.set_inner_cfg_scale(orig_cfg_scale)
|
||||
mr = ModelResult(
|
||||
call_index,
|
||||
sigma,
|
||||
@@ -256,8 +286,9 @@ class ModelCallCache:
|
||||
self.set(call_index, mr)
|
||||
self._fr_add_mr(filter_refs, mr)
|
||||
mr = self.filter_result(mr, default_ref=x, refs=filter_refs)
|
||||
return (mr, False) if return_cached else mr
|
||||
return mr
|
||||
denoised, denoised_prime = torch.func.jvp(call_model, (x, sigma), tangents)
|
||||
self.set_inner_cfg_scale(orig_cfg_scale)
|
||||
mr = ModelResult(
|
||||
call_index,
|
||||
sigma,
|
||||
@@ -270,4 +301,4 @@ class ModelCallCache:
|
||||
self.set(call_index, mr)
|
||||
self._fr_add_mr(filter_refs, mr)
|
||||
mr = self.filter_result(mr, default_ref=x, refs=filter_refs)
|
||||
return (mr, False) if return_cached else mr
|
||||
return mr
|
||||
|
||||
+59
-41
@@ -147,26 +147,15 @@ class SamplerResult:
|
||||
return obj
|
||||
|
||||
|
||||
class CFGPPStepMixin:
|
||||
allow_cfgpp = False
|
||||
allow_alt_cfgpp = False
|
||||
|
||||
def __init__(self):
|
||||
self.cfgpp = self.allow_cfgpp and self.options.pop("cfgpp", False)
|
||||
alt_cfgpp_scale = self.options.pop("alt_cfgpp_scale", 0.0)
|
||||
self.alt_cfgpp_scale = 0.0 if not self.allow_alt_cfgpp else alt_cfgpp_scale
|
||||
|
||||
def to_d(self, mr, **kwargs):
|
||||
return mr.to_d(alt_cfgpp_scale=self.alt_cfgpp_scale, cfgpp=self.cfgpp, **kwargs)
|
||||
|
||||
|
||||
class SingleStepSampler(CFGPPStepMixin):
|
||||
class SingleStepSampler:
|
||||
name = None
|
||||
self_noise = 0
|
||||
model_calls = 0
|
||||
ancestralize = False
|
||||
sample_sigma_zero = False
|
||||
immiscible = None
|
||||
allow_cfgpp = False
|
||||
allow_alt_cfgpp = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -184,7 +173,9 @@ class SingleStepSampler(CFGPPStepMixin):
|
||||
**kwargs,
|
||||
):
|
||||
self.options = kwargs
|
||||
super().__init__()
|
||||
self.cfgpp = self.allow_cfgpp and self.options.pop("cfgpp", False) is True
|
||||
alt_cfgpp_scale = self.options.pop("alt_cfgpp_scale", 0.0)
|
||||
self.alt_cfgpp_scale = 0.0 if not self.allow_alt_cfgpp else alt_cfgpp_scale
|
||||
self.s_noise = s_noise
|
||||
self.eta = eta
|
||||
self.dyn_eta_start = dyn_eta_start
|
||||
@@ -302,9 +293,26 @@ class SingleStepSampler(CFGPPStepMixin):
|
||||
def get_dyn_eta(self, ss):
|
||||
return self.eta * self.get_dyn_value(ss, self.dyn_eta_start, self.dyn_eta_end)
|
||||
|
||||
@property
|
||||
def max_noise_samples(self):
|
||||
return (1 + self.self_noise) * self.substeps
|
||||
|
||||
@property
|
||||
def require_uncond(self):
|
||||
return self.cfgpp or self.alt_cfgpp_scale != 0
|
||||
|
||||
def to_d(self, mr, **kwargs):
|
||||
return mr.to_d(alt_cfgpp_scale=self.alt_cfgpp_scale, cfgpp=self.cfgpp, **kwargs)
|
||||
|
||||
def call_model(self, ss, *args, **kwargs):
|
||||
kwargs["require_uncond"] = self.require_uncond or kwargs.get(
|
||||
"require_uncond", False
|
||||
)
|
||||
kwargs["cfg_scale_override"] = kwargs.get(
|
||||
"cfg_scale_override", ss.cfg_scale_override
|
||||
)
|
||||
return ss.call_model(*args, ss=ss, **kwargs)
|
||||
|
||||
|
||||
class HistorySingleStepSampler(SingleStepSampler):
|
||||
default_history_limit, max_history = 0, 0
|
||||
@@ -383,7 +391,7 @@ class MinSigmaStepMixin:
|
||||
result = yield from self.result(
|
||||
ss, result, sigma_up, sigma=ss.sigma, sigma_next=sn, final=False
|
||||
)
|
||||
mr = ss.model(result, sn, ss=ss, call_index=mcc)
|
||||
mr = self.call_model(ss, result, sn, call_index=mcc)
|
||||
dt = ss.sigma_next - sn
|
||||
result = result + self.to_d(mr) * dt
|
||||
return sigma_up.new_zeros(1), result
|
||||
@@ -542,7 +550,7 @@ class ReversibleHeunStep(ReversibleSingleStepSampler):
|
||||
x_pred = ss.denoised + d * sigma_down
|
||||
|
||||
# Denoised sample at the next sigma
|
||||
mr_next = ss.model(x_pred, sigma_down, ss=ss, call_index=1)
|
||||
mr_next = self.call_model(ss, x_pred, sigma_down, call_index=1)
|
||||
|
||||
# Calculate the derivative at the next sigma
|
||||
d_next = self.to_d(mr_next)
|
||||
@@ -716,7 +724,7 @@ class RESStep(SingleStepSampler):
|
||||
lam_2 = lam + c2_h
|
||||
sigma_2 = lam_2.neg().exp()
|
||||
|
||||
denoised2 = ss.model(x_2, sigma_2, ss=ss, call_index=1).denoised
|
||||
denoised2 = self.call_model(ss, x_2, sigma_2, call_index=1).denoised
|
||||
|
||||
x = math.exp(-h) * eff_x + h * (b1 * denoised + b2 * denoised2)
|
||||
yield from self.result(ss, x, sigma_up, sigma_down=sigma_down)
|
||||
@@ -738,7 +746,7 @@ class TrapezoidalStep(SingleStepSampler):
|
||||
x_pred = x + d_i * ss.dt
|
||||
|
||||
# Denoised sample at the next sigma
|
||||
mr_next = ss.model(x_pred, ss.sigma_next, ss=ss, call_index=1)
|
||||
mr_next = self.call_model(ss, x_pred, ss.sigma_next, call_index=1)
|
||||
|
||||
# Calculate the derivative at the next sigma
|
||||
d_next = self.to_d(mr_next)
|
||||
@@ -762,7 +770,7 @@ class TrapezoidalCycleStep(CycleSingleStepSampler):
|
||||
x_pred = x + d_i * ss.dt
|
||||
|
||||
# Denoised sample at the next sigma
|
||||
mr_next = ss.model(x_pred, ss.sigma_next, ss=ss, call_index=1)
|
||||
mr_next = self.call_model(ss, x_pred, ss.sigma_next, call_index=1)
|
||||
|
||||
# Calculate the derivative at the next sigma
|
||||
d_next = self.to_d(mr_next)
|
||||
@@ -798,10 +806,12 @@ class BogackiStep(ReversibleSingleStepSampler):
|
||||
|
||||
# Bogacki-Shampine steps
|
||||
k1 = d * dt
|
||||
k2 = self.to_d(ss.model(x + k1 / 2, s + dt / 2, ss=ss, call_index=1)) * dt
|
||||
k2 = self.to_d(self.call_model(ss, x + k1 / 2, s + dt / 2, call_index=1)) * dt
|
||||
k3 = (
|
||||
self.to_d(
|
||||
ss.model(x + 3 * k1 / 4 + k2 / 4, s + 3 * dt / 4, ss=ss, call_index=2)
|
||||
self.call_model(
|
||||
ss, x + 3 * k1 / 4 + k2 / 4, s + 3 * dt / 4, call_index=2
|
||||
)
|
||||
)
|
||||
* dt
|
||||
)
|
||||
@@ -834,9 +844,15 @@ class RK4Step(SingleStepSampler):
|
||||
|
||||
# Runge-Kutta steps
|
||||
k1 = d * dt
|
||||
k2 = self.to_d(ss.model(x + k1 / 2, sigma + dt / 2, ss=ss, call_index=1)) * dt
|
||||
k3 = self.to_d(ss.model(x + k2 / 2, sigma + dt / 2, ss=ss, call_index=2)) * dt
|
||||
k4 = self.to_d(ss.model(x + k3, sigma + dt, ss=ss, call_index=3)) * dt
|
||||
k2 = (
|
||||
self.to_d(self.call_model(ss, x + k1 / 2, sigma + dt / 2, call_index=1))
|
||||
* dt
|
||||
)
|
||||
k3 = (
|
||||
self.to_d(self.call_model(ss, x + k2 / 2, sigma + dt / 2, call_index=2))
|
||||
* dt
|
||||
)
|
||||
k4 = self.to_d(self.call_model(ss, x + k3, sigma + dt, call_index=3)) * dt
|
||||
|
||||
# Update the sample
|
||||
x = x + (k1 + 2 * k2 + 2 * k3 + k4) / 6
|
||||
@@ -1087,7 +1103,7 @@ class DPMPP2SStep(SingleStepSampler, DPMPPStepMixin):
|
||||
else x + (ss.denoised - ss.hcur.denoised_uncond) * self.alt_cfgpp_scale
|
||||
)
|
||||
x_2 = (sigma_fn(s) / sigma_fn(t)) * eff_x - (-h * r).expm1() * ss.denoised
|
||||
denoised_2 = ss.model(x_2, sigma_fn(s), ss=ss, call_index=1).denoised
|
||||
denoised_2 = self.call_model(ss, x_2, sigma_fn(s), call_index=1).denoised
|
||||
x = (sigma_fn(t_next) / sigma_fn(t)) * eff_x - (-h).expm1() * denoised_2
|
||||
yield from self.result(ss, x, sigma_up, sigma_down=sigma_down)
|
||||
|
||||
@@ -1123,7 +1139,7 @@ class DPMPPSDEStep(SingleStepSampler, DPMPPStepMixin):
|
||||
x_2 = yield from self.result(
|
||||
ss, x_2, su, sigma=sigma_fn(t), sigma_next=sigma_fn(s), final=False
|
||||
)
|
||||
denoised_2 = ss.model(x_2, sigma_fn(s), ss=ss, call_index=1).denoised
|
||||
denoised_2 = self.call_model(ss, x_2, sigma_fn(s), call_index=1).denoised
|
||||
|
||||
# Step 2
|
||||
sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta)
|
||||
@@ -1154,8 +1170,8 @@ class TTMJVPStep(SingleStepSampler):
|
||||
h_eta = h * (eta + 1)
|
||||
|
||||
eps = to_d(x, sigma, ss.denoised)
|
||||
denoised_prime = ss.model(
|
||||
x, sigma, tangents=(eps * -sigma, -sigma), ss=ss, call_index=1
|
||||
denoised_prime = self.call_model(
|
||||
ss, x, sigma, tangents=(eps * -sigma, -sigma), call_index=1
|
||||
).jdenoised
|
||||
|
||||
phi_1 = -torch.expm1(-h_eta)
|
||||
@@ -1347,7 +1363,7 @@ class HeunPP2Step(SingleStepSampler):
|
||||
w = order * ss.sigma
|
||||
w2 = sn / w
|
||||
x_2 = x + d * dt
|
||||
d_2 = self.to_d(ss.model(x_2, sn, ss=ss, call_index=1))
|
||||
d_2 = self.to_d(self.call_model(ss, x_2, sn, call_index=1))
|
||||
if order == 2:
|
||||
# Heun's method (ish)
|
||||
w1 = 1 - w2
|
||||
@@ -1357,7 +1373,7 @@ class HeunPP2Step(SingleStepSampler):
|
||||
snn = ss.sigmas[ss.idx + 2]
|
||||
dt_2 = snn - sn
|
||||
x_3 = x_2 + d_2 * dt_2
|
||||
d_3 = self.to_d(ss.model(x_3, snn, ss=ss, call_index=2))
|
||||
d_3 = self.to_d(self.call_model(ss, x_3, snn, call_index=2))
|
||||
w3 = snn / w
|
||||
w1 = 1 - w2 - w3
|
||||
d_prime = w1 * d + w2 * d_2 + w3 * d_3
|
||||
@@ -1465,8 +1481,8 @@ class TDEStep(DESolverStep):
|
||||
mcc = 1
|
||||
else:
|
||||
mr_cached = False
|
||||
mr = ss.model(
|
||||
y.unsqueeze(0), t, ss=ss, call_index=mcc, s_in=t.new_ones(1)
|
||||
mr = self.call_model(
|
||||
ss, y.unsqueeze(0), t, call_index=mcc, s_in=t.new_ones(1)
|
||||
)
|
||||
mcc += 1
|
||||
return self.to_d(mr)[bidx if mr_cached else 0]
|
||||
@@ -1578,7 +1594,7 @@ class TODEStep(DESolverStep):
|
||||
mr = ss.hcur
|
||||
mcc = 1
|
||||
else:
|
||||
mr = ss.model(y, t32.clamp(min=1e-05), ss=ss, call_index=mcc)
|
||||
mr = self.call_model(ss, y, t32.clamp(min=1e-05), call_index=mcc)
|
||||
mcc += 1
|
||||
result = self.to_d(mr).flatten(start_dim=1)
|
||||
for bi in range(t.shape[0]):
|
||||
@@ -1708,10 +1724,10 @@ class TSDEStep(DESolverStep):
|
||||
mcc = 1
|
||||
else:
|
||||
mr_cached = False
|
||||
mr = ss.model(
|
||||
mr = outer_self.call_model(
|
||||
ss,
|
||||
y,
|
||||
t32.clamp(min=1e-05),
|
||||
ss=ss,
|
||||
call_index=mcc,
|
||||
s_in=t.new_ones(1),
|
||||
)
|
||||
@@ -1955,13 +1971,15 @@ class DiffraxStep(DESolverStep):
|
||||
mr_cached = False
|
||||
try:
|
||||
if not args:
|
||||
mr = ss.model(y, t32, ss=ss, call_index=mcc, s_in=t.new_ones(1))
|
||||
mr = self.call_model(
|
||||
ss, y, t32, call_index=mcc, s_in=t.new_ones(1)
|
||||
)
|
||||
else:
|
||||
print("TANGENTS")
|
||||
mr = ss.model(
|
||||
mr = self.call_model(
|
||||
ss,
|
||||
y,
|
||||
t32,
|
||||
ss=ss,
|
||||
call_index=mcc,
|
||||
tangents=args,
|
||||
s_in=t.new_ones(1),
|
||||
@@ -2096,7 +2114,7 @@ class HeunStep(ReversibleSingleStepSampler):
|
||||
hcur = ss.hcur
|
||||
d = self.to_d(hcur)
|
||||
x_next = hcur.denoised + d * sd
|
||||
d_next = self.to_d(ss.model(x_next, sd, ss=ss, call_index=1))
|
||||
d_next = self.to_d(self.call_model(ss, x_next, sd, call_index=1))
|
||||
result = hcur.denoised + d * s
|
||||
result += (dt * (d + d_next)) * 0.5
|
||||
result -= self.reversible_correction(ss, d, d_next)
|
||||
@@ -2165,7 +2183,7 @@ class AdapterStep(SingleStepSampler):
|
||||
nonlocal mcc
|
||||
if torch.equal(x_, x) and sigma_ == ss.sigma:
|
||||
return ss.hcur.denoised.clone()
|
||||
mr = ss.model(x_, sigma_, *args, ss=ss, call_index=mcc, **kwargs)
|
||||
mr = self.call_model(ss, x_, sigma_, *args, call_index=mcc, **kwargs)
|
||||
mcc += 1
|
||||
return mr.denoised.clone()
|
||||
|
||||
|
||||
+25
-12
@@ -36,6 +36,8 @@ class MergeSubstepsSampler:
|
||||
self.pre_filter = None if pre_filter is None else make_filter(pre_filter)
|
||||
self.post_filter = None if post_filter is None else make_filter(post_filter)
|
||||
self.preview_mode = options.pop("preview_mode", "denoised")
|
||||
self.require_uncond = any(sampler.require_uncond for sampler in samplers)
|
||||
self.cfg_scale_override = options.pop("cfg_scale_override", None)
|
||||
self.options = options
|
||||
|
||||
def check_match(self, handlers: None | object, *, ss: None | object = None):
|
||||
@@ -107,6 +109,17 @@ class MergeSubstepsSampler:
|
||||
preview_mode = fallback(preview_mode, self.preview_mode)
|
||||
return ss.callback(hi=mr, preview_mode=preview_mode)
|
||||
|
||||
def call_model(self, x, ss=None, sigma=None, **kwargs):
|
||||
ss = fallback(ss, self.ss)
|
||||
sigma = fallback(sigma, ss.sigma)
|
||||
return ss.call_model(
|
||||
x,
|
||||
sigma,
|
||||
ss=ss,
|
||||
cfg_scale_override=self.cfg_scale_override,
|
||||
require_uncond=self.require_uncond,
|
||||
)
|
||||
|
||||
|
||||
class SimpleSubstepsSampler(MergeSubstepsSampler):
|
||||
name = "simple"
|
||||
@@ -130,7 +143,7 @@ class SimpleSubstepsSampler(MergeSubstepsSampler):
|
||||
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
|
||||
)
|
||||
ssampler.noise_sampler = noise_sampler
|
||||
ss.hist.push(ss.model(x, ss.sigma, ss=ss))
|
||||
ss.hist.push(self.call_model(x))
|
||||
ss.refs = FilterRefs.from_ss(ss, have_current=True)
|
||||
self.callback()
|
||||
sr = self.simple_substep(x, ssampler)
|
||||
@@ -149,14 +162,14 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
noise_total = 0.0
|
||||
substep = 0
|
||||
pbar = tqdm.tqdm(total=self.substeps, initial=1, disable=ss.disable_status)
|
||||
ss.hist.push(ss.model(x, ss.sigma, ss=ss))
|
||||
ss.hist.push(self.call_model(x))
|
||||
ss.refs = FilterRefs.from_ss(ss, have_current=True)
|
||||
self.callback()
|
||||
for ssampler in self.samplers:
|
||||
custom_noise = ssampler.custom_noise
|
||||
noise_sampler = ss.noise.make_caching_noise_sampler(
|
||||
custom_noise,
|
||||
ssampler.max_noise_samples(),
|
||||
ssampler.max_noise_samples,
|
||||
ss.sigma,
|
||||
ss.sigma_next,
|
||||
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
|
||||
@@ -226,7 +239,7 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
# noise_sampler = ss.noise.make_caching_noise_sampler(
|
||||
# custom_noise,
|
||||
# ssampler.substeps
|
||||
# + (0 if ss.sigma_next == 0 else ssampler.max_noise_samples()),
|
||||
# + (0 if ss.sigma_next == 0 else ssampler.max_noise_samples),
|
||||
# ss.sigma,
|
||||
# ss.sigma_next,
|
||||
# )
|
||||
@@ -288,7 +301,7 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
# )
|
||||
# noise_sampler = ss.noise.make_caching_noise_sampler(
|
||||
# custom_noise,
|
||||
# ssampler.max_noise_samples() + ssampler.substeps,
|
||||
# ssampler.max_noise_samples + ssampler.substeps,
|
||||
# ss.sigma,
|
||||
# ss.sigma_next,
|
||||
# )
|
||||
@@ -331,7 +344,7 @@ class SupremeAvgMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
# final = merge_ss.sigma_next == 0
|
||||
# noise_sampler = merge_ss.noise.make_caching_noise_sampler(
|
||||
# msampler.options.get("custom_noise", self.options.get("custom_noise")),
|
||||
# msampler.max_noise_samples() + int(not final),
|
||||
# msampler.max_noise_samples + int(not final),
|
||||
# merge_ss.sigma,
|
||||
# merge_ss.sigma_next,
|
||||
# )
|
||||
@@ -396,7 +409,7 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
custom_noise = ssampler.custom_noise
|
||||
noise_sampler = ss.noise.make_caching_noise_sampler(
|
||||
custom_noise,
|
||||
ssampler.max_noise_samples(),
|
||||
ssampler.max_noise_samples,
|
||||
ss.sigma,
|
||||
ss.sigma_next,
|
||||
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
|
||||
@@ -407,7 +420,7 @@ class DivideMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
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.hist.push(self.call_model(x, ss=subss))
|
||||
subss.refs = FilterRefs.from_ss(subss, have_current=True)
|
||||
if substep == 0:
|
||||
self.callback(ss=subss)
|
||||
@@ -476,7 +489,7 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
custom_noise = ssampler.custom_noise
|
||||
noise_sampler = ss.noise.make_caching_noise_sampler(
|
||||
custom_noise,
|
||||
ssampler.max_noise_samples(),
|
||||
ssampler.max_noise_samples,
|
||||
ss.sigma,
|
||||
ss.sigma_next,
|
||||
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
|
||||
@@ -487,7 +500,7 @@ class OvershootMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
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.hist.push(self.call_model(x, ss=subss))
|
||||
subss.refs = FilterRefs.from_ss(subss, have_current=True)
|
||||
if substep == 0:
|
||||
ss.hist.push(subss.hcur)
|
||||
@@ -546,7 +559,7 @@ class LookaheadMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
break
|
||||
noise_sampler = ss.noise.make_caching_noise_sampler(
|
||||
ssampler.custom_noise,
|
||||
ssampler.max_noise_samples(),
|
||||
ssampler.max_noise_samples,
|
||||
ss.sigma,
|
||||
ss.sigma_next,
|
||||
immiscible=fallback(ssampler.immiscible, ss.noise.immiscible),
|
||||
@@ -557,7 +570,7 @@ class LookaheadMergeSubstepsSampler(MergeSubstepsSampler):
|
||||
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.hist.push(self.call_model(x, ss=subss))
|
||||
subss.refs = FilterRefs.from_ss(subss, have_current=True)
|
||||
if substep == 0:
|
||||
self.callback(ss=subss)
|
||||
|
||||
@@ -82,6 +82,7 @@ class StepSamplerGroups(CommonOptionsItems):
|
||||
|
||||
class SamplerState:
|
||||
CLONE_KEYS = (
|
||||
"cfg_scale_override",
|
||||
"model",
|
||||
"hist",
|
||||
"extra_args",
|
||||
@@ -123,6 +124,7 @@ class SamplerState:
|
||||
s_noise=1.0,
|
||||
disable_status=False,
|
||||
history_size=4,
|
||||
cfg_scale_override=None,
|
||||
):
|
||||
self.model = model
|
||||
self.hist = History(max(1, history_size))
|
||||
@@ -138,6 +140,7 @@ class SamplerState:
|
||||
self.step = 0
|
||||
self.substep = 0
|
||||
self.total_steps = len(sigmas) - 1
|
||||
self.cfg_scale_override = cfg_scale_override
|
||||
self.update(idx) # Sets idx, sigma_prev, sigma, sigma_down, refs
|
||||
|
||||
@property
|
||||
@@ -247,3 +250,7 @@ class SamplerState:
|
||||
def reset(self):
|
||||
self.hist.reset()
|
||||
self.denoised = None
|
||||
|
||||
def call_model(self, *args, **kwargs):
|
||||
cfg_scale_override = kwargs.pop("cfg_scale_override", self.cfg_scale_override)
|
||||
return self.model(*args, cfg_scale_override=cfg_scale_override, **kwargs)
|
||||
|
||||
Reference in New Issue
Block a user