support models in nodes / skip configs / input image support mask
This commit is contained in:
@@ -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
@@ -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()))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user