diff --git a/adv_control/sampling.py b/adv_control/sampling.py index 7252160..7f2e9ec 100644 --- a/adv_control/sampling.py +++ b/adv_control/sampling.py @@ -20,12 +20,13 @@ from .control_reference import (ReferenceAdvanced, ReferenceInjections, from .dinklink import get_dinklink from .utils import torch_dfs, WrapperConsts +CURRENT_WRAPPER_VERSION = 10001 def prepare_dinklink_acn_wrapper(): # expose acn_sampler_sample_wrapper d = get_dinklink() link_acn = d.setdefault(WrapperConsts.ACN, {}) - link_acn[WrapperConsts.VERSION] = 10000 + link_acn[WrapperConsts.VERSION] = CURRENT_WRAPPER_VERSION link_acn[WrapperConsts.ACN_CREATE_SAMPLER_SAMPLE_WRAPPER] = (comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY, acn_outer_sample_wrapper) diff --git a/adv_control/utils.py b/adv_control/utils.py index 3b55ff4..bd9629e 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -244,6 +244,11 @@ class TimestepKeyframe: def has_mask_hint(self): return self.mask_hint_orig is not None + def get_effective_guarantee_steps(self, max_sigma: torch.Tensor): + '''If keyframe starts before current sampling range (max_sigma), treat as 0.''' + if self.start_t > max_sigma: + return 0 + return self.guarantee_steps @staticmethod def default() -> 'TimestepKeyframe': @@ -549,7 +554,7 @@ class AdvancedControlBase: self.weights = None self.latent_keyframes = None - def prepare_current_timestep(self, t: Tensor, batched_number: int=1): + def prepare_current_timestep(self, t: Tensor, transformer_options: dict[str, torch.Tensor]): self.t = float(t[0]) # check if t has changed (otherwise do nothing, as step already accounted for) if self.t == self.prev_t: @@ -557,8 +562,9 @@ class AdvancedControlBase: # get current step percent curr_t: float = self.t prev_index = self._current_timestep_index + max_sigma = torch.max(transformer_options.get("sigmas", BIGMAX)) # if met guaranteed steps (or no current keyframe), look for next keyframe in case need to switch - if self._current_timestep_keyframe is None or self._current_used_steps >= self._current_timestep_keyframe.guarantee_steps: + if self._current_timestep_keyframe is None or self._current_used_steps >= self._current_timestep_keyframe.get_effective_guarantee_steps(max_sigma): # if has next index, loop through and see if need to switch if self.timestep_keyframes.has_index(self._current_timestep_index+1): for i in range(self._current_timestep_index+1, len(self.timestep_keyframes)): @@ -584,7 +590,7 @@ class AdvancedControlBase: del self.tk_mask_cond_hint_original self.tk_mask_cond_hint_original = None # if guarantee_steps greater than zero, stop searching for other keyframes - if self._current_timestep_keyframe.guarantee_steps > 0: + if self._current_timestep_keyframe.get_effective_guarantee_steps(max_sigma) > 0: break # if eval_tk is outside of percent range, stop looking further else: @@ -673,7 +679,7 @@ class AdvancedControlBase: self.batch_size = len(t) self.cond_or_uncond = transformer_options.get("cond_or_uncond", None) # prepare timestep and everything related - self.prepare_current_timestep(t=t, batched_number=batched_number) + self.prepare_current_timestep(t=t, transformer_options=transformer_options) # if should not perform any actions for the controlnet, exit without doing any work if self.strength == 0.0 or self._current_timestep_keyframe.strength == 0.0: return self.default_control_actions(x_noisy, t, cond, batched_number, transformer_options) diff --git a/pyproject.toml b/pyproject.toml index e166b84..0477348 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-advanced-controlnet" description = "Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks." -version = "1.4.2" +version = "1.5.0" license = { file = "LICENSE" } dependencies = []