diff --git a/adv_control/control.py b/adv_control/control.py index e31f0e5..a1775e9 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -61,10 +61,11 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): # make cond_hint appropriate dimensions # TODO: change this to not require cond_hint upscaling every step when self.sub_idxs are present - if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * self.compression_ratio != self.cond_hint.shape[2] or x_noisy.shape[3] * self.compression_ratio != self.cond_hint.shape[3]: + if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * self.real_compression_ratio != self.cond_hint.shape[2] or x_noisy.shape[3] * self.real_compression_ratio != self.cond_hint.shape[3]: if self.cond_hint is not None: del self.cond_hint self.cond_hint = None + self.real_compression_ratio = self.compression_ratio compression_ratio = self.compression_ratio if self.vae is not None and self.mult_by_ratio_when_vae: compression_ratio *= self.vae.downscale_ratio @@ -80,6 +81,8 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): loaded_models = comfy.model_management.loaded_models(only_currently_used=True) self.cond_hint = self.vae.encode(self.cond_hint.movedim(1, -1)) comfy.model_management.load_models_gpu(loaded_models) + if not self.mult_by_ratio_when_vae: + self.real_compression_ratio = 1 if self.latent_format is not None: self.cond_hint = self.latent_format.process_in(self.cond_hint) self.cond_hint = self.cond_hint.to(device=x_noisy.device, dtype=dtype) diff --git a/adv_control/utils.py b/adv_control/utils.py index 869cafb..e89ba48 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -504,6 +504,8 @@ class AdvancedControlBase: # vae to store self.adv_vae = None self.mult_by_ratio_when_vae = True + # compression ratio stuff + self.real_compression_ratio = None # require model/vae to be passed into Apply Advanced ControlNet 🛂🅐🅒🅝 node self.require_vae = require_vae self.allow_condhint_latents = allow_condhint_latents @@ -634,6 +636,9 @@ class AdvancedControlBase: # for each timestep keyframe, calculate the start_t for tk in self.timestep_keyframes.keyframes: tk.start_t = percent_to_timestep_function(tk.start_percent) + # set real_compression_ratio to compression_ratio + if hasattr(self, "compression_ratio"): + self.real_compression_ratio = self.compression_ratio # clear variables self.cleanup_advanced() @@ -857,6 +862,9 @@ class AdvancedControlBase: self.batch_size = 0 self.weights = None self.latent_keyframes = None + # set effective_compression_ratio to compression_ratio + if hasattr(self, "compression_ratio"): + self.real_compression_ratio = self.compression_ratio # timestep stuff self._current_timestep_keyframe = None self._current_timestep_index = -1