From 7329728aff75b6b60aee70bd9657c104ff772a92 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 6 Feb 2024 08:21:49 -0600 Subject: [PATCH 1/2] Initial scaffolding for eventual SVD-ControlNet support --- adv_control/control.py | 17 ++++++++++++++++- adv_control/control_svd.py | 6 ++++++ adv_control/utils.py | 1 + 3 files changed, 23 insertions(+), 1 deletion(-) create mode 100644 adv_control/control_svd.py diff --git a/adv_control/control.py b/adv_control/control.py index cd6783b..ee9e3ba 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -168,6 +168,12 @@ class ControlLoraAdvanced(ControlLora, AdvancedControlBase): global_average_pooling=v.global_average_pooling, device=v.device) + +class SVDControlNetAdvanced(ControlNetAdvanced): + def __init__(self, control_model, 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) + + 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) @@ -396,8 +402,9 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo controlnet_type = ControlWeightType.DEFAULT has_controlnet_key = False has_motion_modules_key = False + has_temporal_res_block_key = False for key in controlnet_data: - # LLLLite check + # LLLite check if "lllite" in key: controlnet_type = ControlWeightType.CONTROLLLLITE break @@ -406,14 +413,22 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo has_motion_modules_key = True elif "controlnet" in key: has_controlnet_key = True + # SVD-ControlNet check + elif "temporal_res_block" in key: + has_temporal_res_block_key = True if has_controlnet_key and has_motion_modules_key: controlnet_type = ControlWeightType.SPARSECTRL + elif has_controlnet_key and has_temporal_res_block_key: + controlnet_type = ControlWeightType.SVD_CONTROLNET if controlnet_type != ControlWeightType.DEFAULT: if controlnet_type == ControlWeightType.CONTROLLLLITE: control = load_controllllite(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe) elif controlnet_type == ControlWeightType.SPARSECTRL: control = load_sparsectrl(ckpt_path, controlnet_data=controlnet_data, timestep_keyframe=timestep_keyframe, model=model) + elif controlnet_type == ControlWeightType.SVD_CONTROLNET: + raise Exception(f"SVD-ControlNet is not supported yet!") + #control = comfy_cn.load_controlnet(ckpt_path, model=model) # otherwise, load vanilla ControlNet else: try: diff --git a/adv_control/control_svd.py b/adv_control/control_svd.py new file mode 100644 index 0000000..5832ee7 --- /dev/null +++ b/adv_control/control_svd.py @@ -0,0 +1,6 @@ +from comfy.cldm.cldm import ControlNet as ControlNetCLDM + + +class SVDControlNet(ControlNetCLDM): + def __init__(self, *args,**kwargs): + super().__init__(*args, **kwargs) \ No newline at end of file diff --git a/adv_control/utils.py b/adv_control/utils.py index 3ac0c42..1d77707 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -38,6 +38,7 @@ class ControlWeightType: CONTROLNET = "controlnet" CONTROLLORA = "controllora" CONTROLLLLITE = "controllllite" + SVD_CONTROLNET = "svd_controlnet" SPARSECTRL = "sparsectrl" From 62ce190ec45e984ff3c36f4d7f8a78f9021997f0 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 10 Feb 2024 16:46:24 -0600 Subject: [PATCH 2/2] Added changes from ComfyUI, scaffolding for ReferenceCN support --- adv_control/control.py | 57 ++++++++++++++++++++++++++++++-- adv_control/control_reference.py | 0 2 files changed, 55 insertions(+), 2 deletions(-) create mode 100644 adv_control/control_reference.py diff --git a/adv_control/control.py b/adv_control/control.py index ee9e3ba..d3ebc3b 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -64,7 +64,7 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): # prepare mask_cond_hint self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype) - context = cond['c_crossattn'] + context = cond.get('crossattn_controlnet', cond['c_crossattn']) # uses 'y' in new ComfyUI update y = cond.get('y', None) if y is None: # TODO: remove this in the future since no longer used by newest ComfyUI @@ -288,6 +288,60 @@ class SparseCtrlAdvanced(ControlNetAdvanced): return c +class ReferenceAdvanced(ControlBase, AdvancedControlBase): + def __init__(self, timestep_keyframes: TimestepKeyframeGroup, device=None): + super().__init__(device) + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True) + # TODO: save attn patches here + + def patch_model(self, model: ModelPatcher): + # TODO: do model patching here + pass + + def pre_run_advanced(self, *args, **kwargs): + AdvancedControlBase.pre_run_advanced(self, *args, **kwargs) + # TODO: set control on patches + + def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): + # normal ControlNet stuff + control_prev = None + if self.previous_controlnet is not None: + control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number) + + if self.timestep_range is not None: + if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]: + return control_prev + + dtype = x_noisy.dtype + # prepare cond_hint + if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]: + if self.cond_hint is not None: + del self.cond_hint + self.cond_hint = None + # if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling + 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] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.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) + if x_noisy.shape[0] != self.cond_hint.shape[0]: + self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number) + # prepare mask + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) + # done preparing; model patches will take care of everything now. + # return normal controlnet stuff + return control_prev + + def cleanup_advanced(self): + super().cleanup_advanced() + # TODO: cleanup patches here + + def copy(self): + c = ReferenceAdvanced(self.timestep_keyframes) + self.copy_to(c) + self.copy_to_advanced(c) + return c + + class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): # This ControlNet is more of an attention patch than a traditional controlnet def __init__(self, patch_attn1: LLLitePatch, patch_attn2: LLLitePatch, timestep_keyframes: TimestepKeyframeGroup, device=None): @@ -297,7 +351,6 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): self.patch_attn2 = patch_attn2.set_control(self) self.latent_dims_div2 = None self.latent_dims_div4 = None - def patch_model(self, model: ModelPatcher): model.set_model_attn1_patch(self.patch_attn1) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py new file mode 100644 index 0000000..e69de29