diff --git a/modules/controlnet_nodes.py b/modules/controlnet_nodes.py index 88ab11f..53ed341 100644 --- a/modules/controlnet_nodes.py +++ b/modules/controlnet_nodes.py @@ -95,30 +95,113 @@ except Exception as e: print(e) +def load_controlnet(control_net_name, control_net_override="None"): + if control_net_override != "None": + if control_net_override not in folder_paths.get_filename_list("controlnet"): + print( + f"Warning: Not found ControlNet model {control_net_override}. Use {control_net_name} instead." + ) + else: + control_net_name = control_net_override + + if control_net_name == "None": + return None + + controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) + return comfy.controlnet.load_controlnet(controlnet_path) + + +def apply_preprocessor(image, preprocessor): + if preprocessor == "None": + return image + + if preprocessor not in control_net_preprocessors: + raise Exception(f"Preprocessor {preprocessor} not found") + + preprocessor_class, default_args = control_net_preprocessors[preprocessor] + default_args: List = default_args.copy() + default_args.insert(0, image) + preprocessor_args = { + key: default_args[i] + for i, key in enumerate(preprocessor_class.INPUT_TYPES()["required"].keys()) + } + function_name = preprocessor_class.FUNCTION + image = getattr(preprocessor_class(), function_name)(**preprocessor_args)[0] + + return image + + class AVControlNetLoader(ControlNetLoader): @classmethod def INPUT_TYPES(s): - inputs = ControlNetLoader.INPUT_TYPES() - inputs["optional"] = {"control_net_override": ("STRING", {"default": "None"})} - return inputs + return { + "required": { + "control_net_name": (folder_paths.get_filename_list("controlnet"),) + }, + "optional": {"control_net_override": ("STRING", {"default": "None"})}, + } + RETURN_TYPES = ("CONTROL_NET",) + FUNCTION = "load_controlnet" CATEGORY = "Art Venture/Loaders" def load_controlnet(self, control_net_name, control_net_override="None"): - if control_net_override != "None": - if control_net_override not in folder_paths.get_filename_list("controlnet"): - print( - f"Warning: Not found ControlNet model {control_net_override}. Use {control_net_name} instead." - ) - else: - control_net_name = control_net_override + return load_controlnet(control_net_name, control_net_override) - return super().load_controlnet(control_net_name) + +class AVControlNetEfficientStacker: + controlnets = folder_paths.get_filename_list("controlnet") + preprocessors = list(control_net_preprocessors.keys()) + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "control_net_name": (["None"] + s.controlnets,), + "image": ("IMAGE",), + "strength": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}, + ), + "preprocessor": (["None"] + s.preprocessors,), + }, + "optional": { + "cnet_stack": ("CONTROL_NET_STACK",), + "control_net_override": ("STRING", {"default": "None"}), + }, + } + + RETURN_TYPES = ("CONTROL_NET_STACK",) + RETURN_NAMES = ("CNET_STACK",) + FUNCTION = "control_net_stacker" + CATEGORY = "Art Venture/Loaders" + + def control_net_stacker( + self, + control_net_name, + image, + strength, + preprocessor, + cnet_stack=None, + control_net_override="None", + ): + # If control_net_stack is None, initialize as an empty list + if cnet_stack is None: + cnet_stack = [] + + control_net = load_controlnet(control_net_name, control_net_override) + + # Extend the control_net_stack with the new tuple + if control_net is not None: + image = apply_preprocessor(image, preprocessor) + cnet_stack.extend([(control_net, image, strength)]) + + return (cnet_stack,) class AVControlNetEfficientLoader(ControlNetApply): controlnets = folder_paths.get_filename_list("controlnet") - preprocessors = ["None"] + list(control_net_preprocessors.keys()) + preprocessors = list(control_net_preprocessors.keys()) @classmethod def INPUT_TYPES(s): @@ -131,7 +214,7 @@ class AVControlNetEfficientLoader(ControlNetApply): "FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}, ), - "preprocessor": (s.preprocessors,), + "preprocessor": (["None"] + s.preprocessors,), }, "optional": {"control_net_override": ("STRING", {"default": "None"})}, } @@ -149,32 +232,11 @@ class AVControlNetEfficientLoader(ControlNetApply): preprocessor, control_net_override="None", ): - if control_net_override != "None": - if control_net_override not in self.controlnets: - print( - f"Warning: Not found ControlNet model {control_net_override}. Use {control_net_name} instead." - ) - else: - control_net_name = control_net_override - - if control_net_name == "None": + control_net = load_controlnet(control_net_name, control_net_override) + if control_net is None: return (conditioning,) - controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - control_net = comfy.controlnet.load_controlnet(controlnet_path) - - if preprocessor != "None": - preprocessor_class, default_args = control_net_preprocessors[preprocessor] - default_args: List = default_args.copy() - default_args.insert(0, image) - preprocessor_args = { - key: default_args[i] - for i, key in enumerate( - preprocessor_class.INPUT_TYPES()["required"].keys() - ) - } - function_name = preprocessor_class.FUNCTION - image = getattr(preprocessor_class(), function_name)(**preprocessor_args)[0] + image = apply_preprocessor(image, preprocessor) return super().apply_controlnet(conditioning, control_net, image, strength) @@ -182,9 +244,11 @@ class AVControlNetEfficientLoader(ControlNetApply): NODE_CLASS_MAPPINGS = { "AV_ControlNetLoader": AVControlNetLoader, "AV_ControlNetEfficientLoader": AVControlNetEfficientLoader, + "AV_ControlNetEfficientStacker": AVControlNetEfficientStacker, } NODE_DISPLAY_NAME_MAPPINGS = { "AV_ControlNetLoader": "ControlNet Loader", "AV_ControlNetEfficientLoader": "ControlNet Loader (Efficient)", + "AV_ControlNetEfficientStacker": "ControlNet Stacker (Efficient)", }