From 95e26b9b7906415a3a51e10ded09d6b17f908f8d Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 23 Oct 2024 21:57:39 -0500 Subject: [PATCH 1/2] Made code work with new ComfyUI changes (device var removed from ControlBase) --- adv_control/control.py | 40 ++++++++++++++++---------------- adv_control/control_lllite.py | 4 ++-- adv_control/control_plusplus.py | 8 +++---- adv_control/control_reference.py | 8 +++---- adv_control/utils.py | 2 +- 5 files changed, 31 insertions(+), 31 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index 715d725..d957bd6 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -21,8 +21,8 @@ from .logger import logger class ControlNetAdvanced(ControlNet, AdvancedControlBase): - def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, compression_ratio=8, latent_format=None, device=None, load_device=None, manual_cast_dtype=None, extra_conds=["y"], strength_type=StrengthType.CONSTANT): - super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, compression_ratio=compression_ratio, latent_format=latent_format, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, compression_ratio=8, latent_format=None, load_device=None, manual_cast_dtype=None, extra_conds=["y"], strength_type=StrengthType.CONSTANT): + super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, compression_ratio=compression_ratio, latent_format=latent_format, load_device=load_device, manual_cast_dtype=manual_cast_dtype) AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controlnet()) self.is_flux = False self.x_noisy_shape = None @@ -82,7 +82,7 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): comfy.model_management.load_models_gpu(loaded_models) 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=self.device, dtype=dtype) + self.cond_hint = self.cond_hint.to(device=x_noisy.device, dtype=dtype) if x_noisy.shape[0] != self.cond_hint.shape[0]: self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number) @@ -126,7 +126,7 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): @staticmethod def from_vanilla(v: ControlNet, timestep_keyframe: TimestepKeyframeGroup=None) -> 'ControlNetAdvanced': to_return = ControlNetAdvanced(control_model=v.control_model, timestep_keyframes=timestep_keyframe, - global_average_pooling=v.global_average_pooling, compression_ratio=v.compression_ratio, latent_format=v.latent_format, device=v.device, load_device=v.load_device, + global_average_pooling=v.global_average_pooling, compression_ratio=v.compression_ratio, latent_format=v.latent_format, load_device=v.load_device, manual_cast_dtype=v.manual_cast_dtype) v.copy_to(to_return) return to_return @@ -213,8 +213,8 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase): class ControlLoraAdvanced(ControlLora, AdvancedControlBase): - def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None): - super().__init__(control_weights=control_weights, global_average_pooling=global_average_pooling, device=device) + def __init__(self, control_weights, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False): + super().__init__(control_weights=control_weights, global_average_pooling=global_average_pooling) AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllora()) # use some functions from ControlNetAdvanced self.get_control_advanced = ControlNetAdvanced.get_control_advanced.__get__(self, type(self)) @@ -237,14 +237,14 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase): @staticmethod def from_vanilla(v: ControlLora, timestep_keyframe: TimestepKeyframeGroup=None) -> 'ControlLoraAdvanced': to_return = ControlLoraAdvanced(control_weights=v.control_weights, timestep_keyframes=timestep_keyframe, - global_average_pooling=v.global_average_pooling, device=v.device) + global_average_pooling=v.global_average_pooling) v.copy_to(to_return) return to_return class SVDControlNetAdvanced(ControlNetAdvanced): - def __init__(self, control_model: SVDControlNet, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None): - super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + def __init__(self, control_model: SVDControlNet, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, load_device=None, manual_cast_dtype=None): + super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) def set_cond_hint_inject(self, *args, **kwargs): to_return = super().set_cond_hint_inject(*args, **kwargs) @@ -280,9 +280,9 @@ class SVDControlNetAdvanced(ControlNetAdvanced): actual_cond_hint_orig = self.cond_hint_original if self.cond_hint_original.size(0) < self.full_latent_length: actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length) - self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(x_noisy.device) else: - self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(x_noisy.device) if x_noisy.shape[0] != self.cond_hint.shape[0]: self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number) @@ -311,8 +311,8 @@ class SVDControlNetAdvanced(ControlNetAdvanced): class SparseCtrlAdvanced(ControlNetAdvanced): - def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None): - super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, load_device=None, manual_cast_dtype=None): + super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) self.control_model_wrapped = SparseModelPatcher(self.control_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) self.add_compatible_weight(ControlWeightType.SPARSECTRL) self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints @@ -377,25 +377,25 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.local_sparse_idxs_inverse.remove(actual_i) # sub_cond_hint now contains the hints relevant to current x_noisy if hint_order is None: - sub_cond_hint = self.cond_hint_original[hint_idxs].to(dtype).to(self.device) + sub_cond_hint = self.cond_hint_original[hint_idxs].to(dtype).to(x_noisy.device) else: - sub_cond_hint = self.cond_hint_original[hint_order][hint_idxs].to(dtype).to(self.device) + sub_cond_hint = self.cond_hint_original[hint_order][hint_idxs].to(dtype).to(x_noisy.device) # scale cond_hints to match noisy input if self.control_model.use_simplified_conditioning_embedding: # RGB SparseCtrl; the inputs are latents - use bilinear to avoid blocky artifacts sub_cond_hint = self.model_latent_format.process_in(sub_cond_hint) # multiplies by model scale factor - sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3], x_noisy.shape[2], "nearest-exact", "center").to(dtype).to(self.device) + sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3], x_noisy.shape[2], "nearest-exact", "center").to(dtype).to(x_noisy.device) else: # other SparseCtrl; inputs are typical images - sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + sub_cond_hint = comfy.utils.common_upscale(sub_cond_hint, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(x_noisy.device) # prepare cond_hint (b, c, h ,w) cond_shape = list(sub_cond_hint.shape) cond_shape[0] = len(range_idxs) - self.cond_hint = torch.zeros(cond_shape).to(dtype).to(self.device) + self.cond_hint = torch.zeros(cond_shape).to(dtype).to(x_noisy.device) self.cond_hint[local_idxs] = sub_cond_hint[:] # prepare cond_mask (b, 1, h, w) cond_shape[1] = 1 - cond_mask = torch.zeros(cond_shape).to(dtype).to(self.device) + cond_mask = torch.zeros(cond_shape).to(dtype).to(x_noisy.device) cond_mask[local_idxs] = self.sparse_settings.sparse_mask_mult * self.weights.extras.get(SparseConst.MASK_MULT, 1.0) # combine cond_hint and cond_mask into (b, c+1, h, w) if not self.sparse_settings.merged: @@ -446,7 +446,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.local_sparse_idxs_inverse = None def copy(self): - c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype) + c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.load_device, self.manual_cast_dtype) self.copy_to(c) self.copy_to_advanced(c) return c diff --git a/adv_control/control_lllite.py b/adv_control/control_lllite.py index 4d9dcc1..e219ab3 100644 --- a/adv_control/control_lllite.py +++ b/adv_control/control_lllite.py @@ -340,9 +340,9 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): actual_cond_hint_orig = self.cond_hint_original if self.cond_hint_original.size(0) < self.full_latent_length: actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length) - self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(x_noisy.device) else: - self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(x_noisy.device) if x_noisy.shape[0] != self.cond_hint.shape[0]: self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number) # some special logic here compared to other controlnets: diff --git a/adv_control/control_plusplus.py b/adv_control/control_plusplus.py index cb5c80b..cc2068c 100644 --- a/adv_control/control_plusplus.py +++ b/adv_control/control_plusplus.py @@ -223,8 +223,8 @@ class ControlNetPlusPlus(ControlNetCLDM): class ControlNetPlusPlusAdvanced(ControlNet, AdvancedControlBase): - def __init__(self, control_model: ControlNetPlusPlus, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, device=None, load_device=None, manual_cast_dtype=None): - super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype) + def __init__(self, control_model: ControlNetPlusPlus, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, load_device=None, manual_cast_dtype=None): + super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controlnet()) self.add_compatible_weight(ControlWeightType.CONTROLNETPLUSPLUS) # for IDE type hint purposes @@ -319,13 +319,13 @@ class ControlNetPlusPlusAdvanced(ControlNet, AdvancedControlBase): self.cond_hint[pp_idx] = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, 'nearest-exact', "center") else: self.cond_hint[pp_idx] = comfy.utils.common_upscale(pp_input.image, x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, 'nearest-exact', "center") - self.cond_hint[pp_idx] = self.cond_hint[pp_idx].to(device=self.device, dtype=dtype) + self.cond_hint[pp_idx] = self.cond_hint[pp_idx].to(device=x_noisy.device, dtype=dtype) self.cond_hint_shape = self.cond_hint[pp_idx].shape # prepare cond_hint_controls to match batchsize if self.cond_hint_types.count_nonzero() == 0: self.cond_hint_types = None else: - self.cond_hint_types = self.cond_hint_types.unsqueeze(0).to(device=self.device, dtype=dtype).repeat(x_noisy.shape[0], 1) + self.cond_hint_types = self.cond_hint_types.unsqueeze(0).to(device=x_noisy.device, dtype=dtype).repeat(x_noisy.shape[0], 1) for i in range(len(self.cond_hint)): if self.cond_hint[i] is not None: if x_noisy.shape[0] != self.cond_hint[i].shape[0]: diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 3bea56e..f07581e 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -141,8 +141,8 @@ class ReferencePreprocWrapper(AbstractPreprocWrapper): class ReferenceAdvanced(ControlBase, AdvancedControlBase): CHANNEL_TO_MULT = {320: 1, 640: 2, 1280: 4} - def __init__(self, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup, device=None): - super().__init__(device) + def __init__(self, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup): + super().__init__() AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), allow_condhint_latents=True) # TODO: allow vae_optional to be used instead of preprocessor #require_vae=True @@ -261,11 +261,11 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length: self.cond_hint = comfy.utils.common_upscale( self.cond_hint_original[self.sub_idxs], - x_noisy.shape[3], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(self.device) + x_noisy.shape[3], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(x_noisy.device) else: self.cond_hint = comfy.utils.common_upscale( self.cond_hint_original, - x_noisy.shape[3], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(self.device) + x_noisy.shape[3], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(x_noisy.device) if x_noisy.shape[0] != self.cond_hint.shape[0]: self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number, except_one=False) # noise cond_hint based on sigma (current step) diff --git a/adv_control/utils.py b/adv_control/utils.py index 9c7c1bb..1aa430c 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -931,7 +931,7 @@ class AdvancedControlBase: # default dtype to be same as x_noisy if dtype is None: dtype = x_noisy.dtype - setattr(self, attr_name, out_mask.to(dtype=dtype).to(self.device)) + setattr(self, attr_name, out_mask.to(dtype=dtype).to(x_noisy.device)) del out_mask def _reset_attr(self, attr_name, new_value=None): From 934f55df85b4474e985beef738d203d96295f11a Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 23 Oct 2024 21:58:20 -0500 Subject: [PATCH 2/2] version bump --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 86ccc57..d16181e 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.2.3" +version = "1.3.0" license = { file = "LICENSE" } dependencies = []