From 27613ed685bee03af9928b840571ff5c18f3c5cc Mon Sep 17 00:00:00 2001 From: wl-zhao Date: Fri, 22 Nov 2024 11:49:10 +0800 Subject: [PATCH] support models in nodes / skip configs / input image support mask --- comfy-nodes/input_image.py | 59 ++++++++++++++++++++++++++++++++------ dependency_checker.py | 22 +++++++++++--- 2 files changed, 69 insertions(+), 12 deletions(-) diff --git a/comfy-nodes/input_image.py b/comfy-nodes/input_image.py index b09a645..ebfe847 100755 --- a/comfy-nodes/input_image.py +++ b/comfy-nodes/input_image.py @@ -1,5 +1,7 @@ import folder_paths -from PIL import Image, ImageOps +import node_helpers + +from PIL import Image, ImageOps, ImageSequence, ImageFile import numpy as np import torch import os @@ -32,8 +34,8 @@ class ShellAgentPluginInputImage: } } - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("image",) + RETURN_TYPES = ("IMAGE", "MASK") + # RETURN_NAMES = ("image",) FUNCTION = "run" @@ -60,6 +62,46 @@ class ShellAgentPluginInputImage: return "Invalid image file: {}".format(image) return True + + def convert_image_mask(self, img): + output_images = [] + output_masks = [] + w, h = None, None + + excluded_formats = ['MPO'] + + for i in ImageSequence.Iterator(img): + i = node_helpers.pillow(ImageOps.exif_transpose, i) + + if i.mode == 'I': + i = i.point(lambda i: i * (1 / 255)) + image = i.convert("RGB") + + if len(output_images) == 0: + w = image.size[0] + h = image.size[1] + + if image.size[0] != w or image.size[1] != h: + continue + + image = np.array(image).astype(np.float32) / 255.0 + image = torch.from_numpy(image)[None,] + if 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + else: + mask = torch.zeros((64,64), dtype=torch.float32, device="cpu") + output_images.append(image) + output_masks.append(mask.unsqueeze(0)) + + if len(output_images) > 1 and img.format not in excluded_formats: + output_image = torch.cat(output_images, dim=0) + output_mask = torch.cat(output_masks, dim=0) + else: + output_image = output_images[0] + output_mask = output_masks[0] + + return (output_image, output_mask) def run(self, input_name, default_value=None, display_name=None, description=None): input_dir = folder_paths.get_input_directory() @@ -84,11 +126,12 @@ class ShellAgentPluginInputImage: image_path = os.path.join(input_dir, image_path) image = Image.open(image_path).convert("RGB") - image = ImageOps.exif_transpose(image) - image = image.convert("RGB") - image = np.array(image).astype(np.float32) / 255.0 - image = torch.from_numpy(image)[None,] - return [image] + return self.convert_image_mask(image) + # image = ImageOps.exif_transpose(image) + # image = image.convert("RGB") + # image = np.array(image).astype(np.float32) / 255.0 + # image = torch.from_numpy(image)[None,] + # return [image] except Exception as e: raise e diff --git a/dependency_checker.py b/dependency_checker.py index f2b2200..1e73e05 100644 --- a/dependency_checker.py +++ b/dependency_checker.py @@ -151,8 +151,8 @@ def resolve_dependencies(prompt, custom_dependencies): # resolve custom nodes an file_mapping_dict = {} - - def collect_unknown_models(filename, node_id, node_info): + SKIP_FOLDER_NAMES = ["configs", "custom_nodes"] + def collect_unknown_models(filename, node_id, node_info, custom_node_path): if type(filename) != str: return is_model = False @@ -166,6 +166,9 @@ def resolve_dependencies(prompt, custom_dependencies): # resolve custom nodes an # Walk through all subdirectories and files in the directory rel_save_path = None for possible_folder_name in folder_paths.folder_names_and_paths: + if possible_folder_name in SKIP_FOLDER_NAMES: + print(f"skip {possible_folder_name}") + continue full_path = folder_paths.get_full_path(possible_folder_name, filename) if full_path is None: continue @@ -188,6 +191,16 @@ def resolve_dependencies(prompt, custom_dependencies): # resolve custom nodes an print(f"matched files: {matching_files}") + # step 3: search inside the custom nodes + if custom_node_path is not None: + for full_path in glob.glob(f"{custom_node_path}/**/*", recursive=True): + if os.path.isfile(full_path) and full_path.endswith(filename) and full_path not in matching_files: + folder_path = full_path[:-len(filename)] + rel_save_path = os.path.relpath(folder_path, folder_paths.models_dir) + matching_files[full_path] = { + "rel_save_path": rel_save_path + } + if len(matching_files) == 0: raise ValueError(f"Cannot find model: `{filename}`, Node ID: `{node_id}`, Node Info: `{node_info}`") @@ -210,10 +223,11 @@ def resolve_dependencies(prompt, custom_dependencies): # resolve custom nodes an skip_model_check = False + custom_node_path = None if hasattr(node_cls, "RELATIVE_PYTHON_MODULE") and node_cls.RELATIVE_PYTHON_MODULE.startswith("custom_nodes."): print(node_cls.RELATIVE_PYTHON_MODULE) custom_nodes.append(node_cls.RELATIVE_PYTHON_MODULE) - + custom_node_path = os.path.join(BASE_PATH, node_cls.RELATIVE_PYTHON_MODULE.replace(".", "/")) if node_cls.RELATIVE_PYTHON_MODULE[len("custom_nodes."):] in node_remote_skip_models: skip_model_check = True print(f"skip model check for {node_class_type}") @@ -236,7 +250,7 @@ def resolve_dependencies(prompt, custom_dependencies): # resolve custom nodes an "rel_save_path": rel_save_path } elif not skip_model_check: - tree_map(lambda x: collect_unknown_models(x, node_id, node_info), node_info["inputs"]) + tree_map(lambda x: collect_unknown_models(x, node_id, node_info, custom_node_path), node_info["inputs"]) list(map(partial(collect_local_file, mapping_dict=file_mapping_dict), node_info["inputs"].values()))