From cb773cd8518c67921aaa3a0e4a123d7c3154630f Mon Sep 17 00:00:00 2001 From: blepping Date: Sat, 7 Sep 2024 06:35:52 -0600 Subject: [PATCH] Allow overriding CFG scale in parameters Try to make sure uncond gets generated when needed --- py/filtering.py | 2 + py/model.py | 89 ++++++++++++++++++++++++------------ py/step_samplers.py | 100 ++++++++++++++++++++++++----------------- py/substep_merging.py | 37 ++++++++++----- py/substep_sampling.py | 7 +++ 5 files changed, 153 insertions(+), 82 deletions(-) diff --git a/py/filtering.py b/py/filtering.py index 55488d7..6d62de1 100644 --- a/py/filtering.py +++ b/py/filtering.py @@ -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) diff --git a/py/model.py b/py/model.py index 6bc1111..bd05435 100644 --- a/py/model.py +++ b/py/model.py @@ -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 diff --git a/py/step_samplers.py b/py/step_samplers.py index 2c102b2..0ef7a2e 100644 --- a/py/step_samplers.py +++ b/py/step_samplers.py @@ -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() diff --git a/py/substep_merging.py b/py/substep_merging.py index 8f3cf49..9f4a0db 100644 --- a/py/substep_merging.py +++ b/py/substep_merging.py @@ -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) diff --git a/py/substep_sampling.py b/py/substep_sampling.py index 7c9608f..d003974 100644 --- a/py/substep_sampling.py +++ b/py/substep_sampling.py @@ -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)