Allow overriding CFG scale in parameters

Try to make sure uncond gets generated when needed
This commit is contained in:
blepping
2024-09-07 06:35:52 -06:00
parent 67a479a65d
commit cb773cd851
5 changed files with 153 additions and 82 deletions
+2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+7
View File
@@ -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)