diff --git a/control/control.py b/control/control.py index c9bc138..df53914 100644 --- a/control/control.py +++ b/control/control.py @@ -374,7 +374,7 @@ def is_advanced_controlnet(input_object): return hasattr(input_object, "sub_idxs") -def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, sparse_settings=None, model=None) -> SparseCtrlAdvanced: +def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None, sparse_settings=SparseSettings.default(), model=None) -> SparseCtrlAdvanced: if controlnet_data is None: controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True) # first, separate out motion part from normal controlnet part and attempt to load that portion @@ -512,8 +512,9 @@ def load_sparsectrl(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, tim if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"): #TODO: smarter way of enabling global_average_pooling global_average_pooling = True - # both motion portion and controlnet portions are loaded; bring them together - motion_wrapper.inject(control_model) + # both motion portion and controlnet portions are loaded; bring them together if using motion model + if sparse_settings.use_motion: + motion_wrapper.inject(control_model) control = SparseCtrlAdvanced(control_model, timestep_keyframes=timestep_keyframe, sparse_settings=sparse_settings, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) return control diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 680a6d3..717e745 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -80,12 +80,13 @@ class SparseControlNet(ControlNetCLDM): class SparseSettings: - def __init__(self, sparse_method: 'SparseMethod'): + def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True): self.sparse_method = sparse_method + self.use_motion = use_motion @classmethod def default(cls): - return cls(sparse_method=SparseSpreadMethod()) + return SparseSettings(sparse_method=SparseSpreadMethod(), use_motion=True) class SparseMethod(ABC): diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index e396550..7759013 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -13,6 +13,7 @@ class SparseCtrlLoaderAdvanced: return { "required": { "control_net_name": (folder_paths.get_filename_list("controlnet"), ), + "use_motion": ("BOOLEAN", {"default": True}, ), }, "optional": { "sparse_method": ("SPARSE_METHOD", ), @@ -25,9 +26,9 @@ class SparseCtrlLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def load_controlnet(self, control_net_name: str, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name: str, use_motion: bool, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - sparse_settings = SparseSettings(sparse_method=sparse_method) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion) controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) return (controlnet,) @@ -70,7 +71,7 @@ class SparseSpreadMethodNode: def INPUT_TYPES(s): return { "required": { - "spread": (SparseSpreadMethod.LIST), + "spread": (SparseSpreadMethod.LIST,), } }