From 8086ad3345fd2bdb1672a88bf15b05a19bd52b1d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E9=9B=AA=E5=B3=B0?= Date: Fri, 18 Oct 2024 20:00:58 +0800 Subject: [PATCH] add some nodes for batch --- README.md | 12 ++- easyapi/ImageNode.py | 112 ++++++++++++++++++++++- easyapi/UtilNode.py | 213 ++++++++++++++++++++++++++++++++++++++++++- pyproject.toml | 2 +- 4 files changed, 333 insertions(+), 6 deletions(-) diff --git a/README.md b/README.md index 584a9df..d8f116e 100644 --- a/README.md +++ b/README.md @@ -65,7 +65,14 @@ Tips: base64格式字符串比较长,会导致界面卡顿,接口请求带 | ForEachClose | 循环结束节点 | | LoadJsonStrToList | json字符串转换为对象列表 | | GetValueFromJsonObj | 从对象中获取指定key的值 | -| FilterValueForList | 根据指定值过滤列表中元素 || +| FilterValueForList | 根据指定值过滤列表中元素 | +| SliceList | 列表切片 | +| LoadLocalFilePath | 列出给定路径下的文件列表 | +| LoadImageFromLocalPath | 根据图片全路径加载图片 | +| LoadMaskFromLocalPath | 根据遮罩全路径加载遮罩 | | +| IsNoneOrEmpty | 判断是否为空或空字符串或空列表或空字典 | +| IsNoneOrEmptyOptional | 为空时返回指定值(惰性求值),否则返回原值 | +| EmptyOutputNode | 空的输出类型节点 | ### 示例 ![save api extended](docs/example_note.png) @@ -76,6 +83,9 @@ Tips: base64格式字符串比较长,会导致界面卡顿,接口请求带 ![save api extended](example/example_3.png) ## 更新记录 +### 2024-10-18 +- 新增节点:SliceList、LoadLocalFilePath、LoadImageFromLocalPath、LoadMaskFromLocalPath、IsNoneOrEmpty、IsNoneOrEmptyOptional、EmptyOutputNode + ### 2024-09-29 - 新增节点:FilterValueForList diff --git a/easyapi/ImageNode.py b/easyapi/ImageNode.py index 2c30579..c75eaa5 100644 --- a/easyapi/ImageNode.py +++ b/easyapi/ImageNode.py @@ -1,9 +1,13 @@ import base64 import copy import io +import os + import numpy as np import torch -from PIL import ImageOps, Image +from PIL import ImageOps, Image, ImageSequence + +import node_helpers from nodes import LoadImage from comfy.cli_args import args from PIL.PngImagePlugin import PngInfo @@ -340,6 +344,108 @@ class LoadImageToBase64(LoadImage): return encoded_image, img, mask +class LoadImageFromLocalPath: + @classmethod + def INPUT_TYPES(s): + return {"required": + { + "image_path": ("STRING", {"default": ""},) + }, + } + + CATEGORY = "EasyApi/Image" + + RETURN_TYPES = ("IMAGE", "MASK") + FUNCTION = "load_image" + def load_image(self, image_path): + + img = node_helpers.pillow(Image.open, image_path) + + 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)) + # 将图像转换为RGB格式 + 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 + + # 将图像转换为浮点数组 (H,W,Channel) + image = np.array(image).astype(np.float32) / 255.0 + # 先把图片转成3维张量,并再在最前面添加一个维度,变成4维(1, H, W,Channel) + image = torch.from_numpy(image)[None,] + # 如果图像包含alpha通道,则将其转换为掩码 + if 'A' in i.getbands(): + # 计算后结果数组中透明像素会是0 + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + # 把数组中透明像素设为1 + mask = 1. - torch.from_numpy(mask) + else: + # 否则,创建一个64x64的零张量作为掩码 + 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: + # 如果有多个图像,则将它们按维度0拼接在一起 + 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) + + +class LoadMaskFromLocalPath: + _color_channels = ["alpha", "red", "green", "blue"] + @classmethod + def INPUT_TYPES(s): + return {"required": + { + "image_path": ("STRING", {"default": ""}), + "channel": (s._color_channels, ), + } + } + + CATEGORY = "EasyApi/Image" + + RETURN_TYPES = ("MASK",) + FUNCTION = "load_mask" + def load_mask(self, image_path, channel): + i = node_helpers.pillow(Image.open, image_path) + i = node_helpers.pillow(ImageOps.exif_transpose, i) + if i.getbands() != ("R", "G", "B", "A"): + if i.mode == 'I': + i = i.point(lambda i: i * (1 / 255)) + i = i.convert("RGBA") + mask = None + c = channel[0].upper() + if c in i.getbands(): + mask = np.array(i.getchannel(c)).astype(np.float32) / 255.0 + mask = torch.from_numpy(mask) + if c == 'A': + mask = 1. - mask + else: + mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu") + return (mask.unsqueeze(0),) + + NODE_CLASS_MAPPINGS = { "Base64ToImage": Base64ToImage, "LoadImageFromURL": LoadImageFromURL, @@ -351,6 +457,8 @@ NODE_CLASS_MAPPINGS = { "MaskToBase64Image": MaskToBase64Image, "MaskImageToBase64": MaskImageToBase64, "LoadImageToBase64": LoadImageToBase64, + "LoadImageFromLocalPath": LoadImageFromLocalPath, + "LoadMaskFromLocalPath": LoadMaskFromLocalPath, } # A dictionary that contains the friendly/humanly readable titles for the nodes @@ -365,4 +473,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "MaskToBase64Image": "Mask To Base64 Image", "MaskImageToBase64": "Mask Image To Base64", "LoadImageToBase64": "Load Image To Base64", + "LoadImageFromLocalPath": "Load Image From Local Path", + "LoadMaskFromLocalPath": "Load Mask From Local Path", } diff --git a/easyapi/UtilNode.py b/easyapi/UtilNode.py index 29634f7..87f7910 100644 --- a/easyapi/UtilNode.py +++ b/easyapi/UtilNode.py @@ -1,6 +1,10 @@ +import mimetypes +import os + import simplejson import torch +import folder_paths from comfy.model_patcher import ModelPatcher import comfy.model_base from .util import tensor_to_pil, hex_to_rgba, any_type @@ -136,6 +140,9 @@ class SplitStringToList: "str": ('STRING', {"forceInput": True}), "to_type": (["str", "int", "float", "bool"], {"default": "str"}), "delimiter": ('STRING', {"default": ","}), + }, + "optional": { + "method": (["delimiter", "LF", "tab"], {"default": "delimiter", "tooltip": "分隔符选取方式"}), } } @@ -147,7 +154,11 @@ class SplitStringToList: CATEGORY = "EasyApi/String" DESCRIPTION = "按分隔符把字符串拆分成列表。如 \"a,b,c\" => [a,b,c]" - def convert(self, str, to_type, delimiter): + def convert(self, str, to_type, delimiter, method="delimiter"): + if method == "LF": + delimiter = "\n" + elif method == "tab": + delimiter = "\t" result = [item.strip() for item in str.split(delimiter)] if to_type == "int": result = [int(x) for x in result] @@ -458,7 +469,7 @@ class IndexOfList: return { "required": { "lst": (any_type, {}), - "index": ('INT', {'default': 0, 'step': 1, 'min': 0, 'max': 50}), + "index": ('INT', {'default': 0, 'step': 1, 'min': 0, 'max': 100000}), } } @@ -469,7 +480,7 @@ class IndexOfList: CATEGORY = "EasyApi/List" - DESCRIPTION = "根据索引过滤" + DESCRIPTION = "根据索引过滤,若index >= len(lst),返回None" def execute(self, lst, index): if isinstance(lst, list) and len(lst) > index: @@ -504,6 +515,37 @@ class IndexesOfList: return (None, ) +class SliceList: + @classmethod + def INPUT_TYPES(self): + return { + "required": { + "lst": (any_type, {}), + "start_index": ('INT', {'default': 0, 'step': 1, 'min': -100000, 'max': 100000}), + "step": ('INT', {'default': 1, 'step': 1, 'min': -100000, 'max': 100000}), + "end_index": ('INT', {'default': 100000, 'step': 1, 'min': -100000, 'max': 100000}), + "reverse": ('BOOLEAN', {'default': False}), + } + } + + RETURN_TYPES = (any_type,) + RETURN_NAMES = ("lst",) + + FUNCTION = "execute" + + CATEGORY = "EasyApi/List" + + DESCRIPTION = "列表切片, lst入参不是list时,返回None" + + def execute(self, lst, start_index, step, end_index, reverse): + if isinstance(lst, list): + sliceList = lst[start_index:end_index:step] + if reverse: + sliceList.reverse() + return (sliceList, ) + return (None, ) + + class StringArea: @classmethod def INPUT_TYPES(s): @@ -625,6 +667,161 @@ class FilterValueForList: return (filtered,) +class LoadLocalFilePath: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "directory": ("STRING", {"default": "", "tooltip": "若为空,遍历input目录"}), + "max_depth": ("INT", {"default": 1, "min": 1, "max": 64, "step": 1, "tooltip": "查找最大目录层级"}), + "file_type": (["image", "video", "text"], {"default": "image", "tooltip": "file_suffix值不为空时,此配置失效"}), + "file_suffix": ("STRING", {"default": "", "tooltip": "指定过滤文件后缀,多个以|分割,如.png|.jpg"}), + } + } + + RETURN_TYPES = ("LIST", "INT",) + RETURN_NAMES = ("paths", "count",) + OUTPUT_TOOLTIPS = ("文件路径列表,若过滤不到文件返回空列表", "文件个数",) + + FUNCTION = "get_paths" + + CATEGORY = "EasyApi/Utils" + + DESCRIPTION = "根据条件遍历指定目录下文件路径" + + mime_types_dict = { + 'image': {'image/jpeg', 'image/png', 'image/gif', 'image/bmp', 'image/tiff'}, + 'video': {'video/mp4', 'video/quicktime', 'video/x-msvideo', 'video/x-matroska'}, + 'text': {'text/plain', 'text/html', 'text/css', 'text/csv'} + } + + @classmethod + def recursive_file_paths(cls, directory, max_depth, file_type, file_suffix, current_depth=1): + """ + 获取指定目录及其子目录中的图片文件路径(深度优先遍历) + + 参数: + directory (str): 要遍历的目录路径 + max_depth (int): 最大遍历层级 + current_depth (int): 当前遍历层级(默认值为1) + + 返回: + List[str]: 图片文件路径列表 + """ + + image_paths = [] + + if current_depth > max_depth: + return image_paths + + with os.scandir(directory) as it: + for item in it: + if item.is_file(): + if len(file_suffix.strip()) > 0: + suffixes = [s.strip().lower() for s in file_suffix.split('|')] + if any(item.name.lower().endswith(suffix) for suffix in suffixes): + image_paths.append(item.path) + elif file_type: + mime_type, _ = mimetypes.guess_type(item.path) + if mime_type in cls.mime_types_dict.get(file_type, set()): + image_paths.append(item.path) + elif item.is_dir(): + image_paths.extend(cls.recursive_file_paths(item.path, max_depth, file_type, file_suffix, current_depth + 1)) + return image_paths + + def get_paths(self, directory, max_depth, file_type, file_suffix): + if directory is None or len(directory.strip()) == 0: + directory = folder_paths.get_input_directory() + image_paths = self.recursive_file_paths(directory, max_depth, file_type, file_suffix) + + return image_paths, len(image_paths), + + +class IsNoneOrEmpty: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "any": (any_type,) + } + } + + RETURN_TYPES = ("BOOLEAN",) + RETURN_NAMES = ("boolean",) + FUNCTION = "execute" + CATEGORY = "EasyApi/Utils" + DESCRIPTION = "判断输入是否为None、空列表、空字符串(trim后判断)、空字典" + + def execute(self, any): + if any is None: + return True, + if isinstance(any, list): + return (True if len(any) == 0 else False,) + if isinstance(any, str): + return (True if len(any.strip()) == 0 else False,) + if isinstance(any, dict): + return (True if len(any) == 0 else False,) + return False + + +class IsNoneOrEmptyOptional: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "any": (any_type,) + }, + "optional": { + "default": (any_type, {"lazy": True}) + } + } + + RETURN_TYPES = (any_type,) + RETURN_NAMES = ("any",) + FUNCTION = "execute" + CATEGORY = "EasyApi/Utils" + DESCRIPTION = "判断输入any是否为None、空列表、空字符串(trim后判断)、空字典,若为true,返回default的值,否则返回输入值" + + def execute(self, any, default=None): + if any is None: + return default, + if isinstance(any, list): + return (default if len(any) == 0 else any,) + if isinstance(any, str): + return (default if len(any.strip()) == 0 else any,) + if isinstance(any, dict): + return (default if len(any) == 0 else any,) + return (any,) + + def check_lazy_status(self, any, default=None): + if any is None: + return ["default"] + if isinstance(any, list): + return ["default"] if len(any) == 0 else ["any"] + if isinstance(any, str): + return ["default"] if len(any.strip()) == 0 else ["any"] + if isinstance(any, dict): + return ["default"] if len(any) == 0 else ["any"] + return ["any"] + + +class EmptyOutputNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "any": (any_type,) + } + } + RETURN_TYPES = () + FUNCTION = "execute" + CATEGORY = "EasyApi/Utils" + DESCRIPTION = "可配合for循环批量处理图片,for循环后连接此输出节点" + OUTPUT_NODE = True + def execute(self, any): + return () + + NODE_CLASS_MAPPINGS = { "GetImageBatchSize": GetImageBatchSize, "JoinList": JoinList, @@ -645,11 +842,16 @@ NODE_CLASS_MAPPINGS = { "SplitStringToList": SplitStringToList, "IndexOfList": IndexOfList, "IndexesOfList": IndexesOfList, + "SliceList": SliceList, "StringArea": StringArea, "ConvertTypeToAny": ConvertTypeToAny, "GetValueFromJsonObj": GetValueFromJsonObj, "LoadJsonStrToList": LoadJsonStrToList, "FilterValueForList": FilterValueForList, + "LoadLocalFilePath": LoadLocalFilePath, + "IsNoneOrEmpty": IsNoneOrEmpty, + "IsNoneOrEmptyOptional": IsNoneOrEmptyOptional, + "EmptyOutputNode": EmptyOutputNode, } # A dictionary that contains the friendly/humanly readable titles for the nodes @@ -673,9 +875,14 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SplitStringToList": "SplitStringToList", "IndexOfList": "IndexOfList", "IndexesOfList": "IndexesOfList", + "SliceList": "SliceList", "StringArea": "StringArea", "ConvertTypeToAny": "ConvertTypeToAny", "GetValueFromJsonObj": "GetValueFromJsonObj", "LoadJsonStrToList": "LoadJsonStrToList", "FilterValueForList": "FilterValueForList", + "LoadLocalFilePath": "LoadLocalFilePath", + "IsNoneOrEmpty": "IsNoneOrEmpty", + "IsNoneOrEmptyOptional": "IsNoneOrEmptyOptional", + "EmptyOutputNode": "EmptyOutputNode", } diff --git a/pyproject.toml b/pyproject.toml index 003e60c..4c5ae33 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-easyapi-nodes" description = "Provides some features and nodes related to API calls." -version = "1.0.5" +version = "1.0.6" license = { file = "LICENSE" } dependencies = ["segment_anything", "simple_lama_inpainting", "insightface"]