Add experimental Load Merged SparseCtrl Model node
This commit is contained in:
+2
-1
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
+3
-1
@@ -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 🛂🅐🅒🅝",
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user