From 7329728aff75b6b60aee70bd9657c104ff772a92 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 6 Feb 2024 08:21:49 -0600 Subject: [PATCH] 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"