Add real_compression_ratio so can be adjusted to properly check for expected cond_hint size

This commit is contained in:
Jedrzej Kosinski
2024-11-25 21:52:51 -06:00
parent 72f0e6dbac
commit 999763ce69
2 changed files with 12 additions and 1 deletions
+4 -1
View File
@@ -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)
+8
View File
@@ -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