Fixed Sparse Spread Method dropdown not rendering, added use_motion toggle

This commit is contained in:
Jedrzej Kosinski
2023-12-21 01:24:13 -06:00
parent be0460ec04
commit 2fc43f458b
3 changed files with 11 additions and 8 deletions
+4 -3
View File
@@ -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
+3 -2
View File
@@ -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):
+4 -3
View File
@@ -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,),
}
}