diff --git a/control/control.py b/control/control.py index b473b3e..32b4912 100644 --- a/control/control.py +++ b/control/control.py @@ -286,7 +286,8 @@ class SparseCtrlAdvanced(ControlNetAdvanced): cond_mask = torch.zeros(cond_shape).to(dtype).to(self.device) cond_mask[local_idxs] = 1.0 # combine cond_hint and cond_mask into (b, c+1, h, w) - self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1) + if not self.sparse_settings.merged: + self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1) del sub_cond_hint del cond_mask # make cond_hint match x_noisy batch diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 1abb8cd..5328cac 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -80,11 +80,12 @@ class SparseControlNet(ControlNetCLDM): class SparseSettings: - def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0): + def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0, merged=False): self.sparse_method = sparse_method self.use_motion = use_motion self.motion_strength = motion_strength self.motion_scale = motion_scale + self.merged = merged @classmethod def default(cls): diff --git a/control/nodes.py b/control/nodes.py index 2ed68e1..5076062 100644 --- a/control/nodes.py +++ b/control/nodes.py @@ -9,7 +9,7 @@ from .utils import StrengthInterpolation as SI from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights, SoftT2IAdapterWeights, CustomT2IAdapterWeights) from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode -from .nodes_sparsectrl import SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, VAEEncodePreprocessor +from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, VAEEncodePreprocessor from .logger import logger @@ -217,6 +217,7 @@ NODE_CLASS_MAPPINGS = { # SparseCtrl "ACN_VAEEncodePreprocessor": VAEEncodePreprocessor, "ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced, + "ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced, "ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode, "ACN_SparseCtrlSpreadMethodNode": SparseSpreadMethodNode, } @@ -244,6 +245,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { # SparseCtrl "ACN_VAEEncodePreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝", "ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝", + "ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝", "ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝", "ACN_SparseCtrlSpreadMethodNode": "SparseCtrl Spread Method 🛂🅐🅒🅝", } diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 2faae66..8a724a7 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -3,7 +3,7 @@ from nodes import VAEEncode from .utils import TimestepKeyframeGroup from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod -from .control import load_sparsectrl +from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced # node for SparseCtrl loading @@ -12,6 +12,35 @@ class SparseCtrlLoaderAdvanced: def INPUT_TYPES(s): return { "required": { + "sparsectrl_name": (folder_paths.get_filename_list("controlnet"), ), + "use_motion": ("BOOLEAN", {"default": True}, ), + "motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + }, + "optional": { + "sparse_method": ("SPARSE_METHOD", ), + "tk_optional": ("TIMESTEP_KEYFRAME", ), + } + } + + RETURN_TYPES = ("CONTROL_NET", ) + FUNCTION = "load_controlnet" + + CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" + + def load_controlnet(self, sparsectrl_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale) + sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) + return (sparsectrl,) + + +class SparseCtrlMergedLoaderAdvanced: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "sparsectrl_name": (folder_paths.get_filename_list("controlnet"), ), "control_net_name": (folder_paths.get_filename_list("controlnet"), ), "use_motion": ("BOOLEAN", {"default": True}, ), "motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), @@ -28,11 +57,24 @@ class SparseCtrlLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def load_controlnet(self, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + def load_controlnet(self, sparsectrl_name: str, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name) controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale) - controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) - return (controlnet,) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale, merged=True) + # first, load normal controlnet + controlnet = load_controlnet(controlnet_path, timestep_keyframe=tk_optional) + # confirm that controlnet is ControlNetAdvanced + if controlnet is None or type(controlnet) != ControlNetAdvanced: + raise ValueError(f"controlnet_path must point to a normal ControlNet, but instead: {type(controlnet).__name__}") + # next, load sparsectrl, making sure to load motion portion + sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=SparseSettings.default()) + # now, combine state dicts + new_state_dict = controlnet.control_model.state_dict() + for key, value in sparsectrl.control_model.motion_holder.motion_wrapper.state_dict().items(): + new_state_dict[key] = value + # now, reload sparsectrl with real settings + sparsectrl = load_sparsectrl(sparsectrl_path, controlnet_data=new_state_dict, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) + return (sparsectrl,) class SparseIndexMethodNode: