support models in nodes / skip configs / input image support mask

This commit is contained in:
wl-zhao
2024-11-22 11:49:10 +08:00
parent 92673900c5
commit 27613ed685
2 changed files with 69 additions and 12 deletions
+51 -8
View File
@@ -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
+18 -4
View File
@@ -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()))