From 1c3e0ca8e4d6609f6cdd34084fe8d9506f3a919c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=94=A1=E5=AD=9F=E6=98=86?= <865240848@qq.com> Date: Sun, 4 Feb 2024 15:13:33 +0800 Subject: [PATCH] fix --- .idea/.gitignore | 8 ++ .idea/comfyui-segment-anything-marko.iml | 12 +++ .../inspectionProfiles/profiles_settings.xml | 6 ++ .idea/misc.xml | 7 ++ .idea/modules.xml | 8 ++ __init__.py | 12 +++ install.py | 17 +++++ node.py | 76 +++++++++++++++++++ requirements.txt | 5 ++ 9 files changed, 151 insertions(+) create mode 100644 .idea/.gitignore create mode 100644 .idea/comfyui-segment-anything-marko.iml create mode 100644 .idea/inspectionProfiles/profiles_settings.xml create mode 100644 .idea/misc.xml create mode 100644 .idea/modules.xml create mode 100644 __init__.py create mode 100644 install.py create mode 100644 node.py create mode 100644 requirements.txt diff --git a/.idea/.gitignore b/.idea/.gitignore new file mode 100644 index 0000000..35410ca --- /dev/null +++ b/.idea/.gitignore @@ -0,0 +1,8 @@ +# 默认忽略的文件 +/shelf/ +/workspace.xml +# 基于编辑器的 HTTP 客户端请求 +/httpRequests/ +# Datasource local storage ignored files +/dataSources/ +/dataSources.local.xml diff --git a/.idea/comfyui-segment-anything-marko.iml b/.idea/comfyui-segment-anything-marko.iml new file mode 100644 index 0000000..039314d --- /dev/null +++ b/.idea/comfyui-segment-anything-marko.iml @@ -0,0 +1,12 @@ + + + + + + + + + + \ No newline at end of file diff --git a/.idea/inspectionProfiles/profiles_settings.xml b/.idea/inspectionProfiles/profiles_settings.xml new file mode 100644 index 0000000..105ce2d --- /dev/null +++ b/.idea/inspectionProfiles/profiles_settings.xml @@ -0,0 +1,6 @@ + + + + \ No newline at end of file diff --git a/.idea/misc.xml b/.idea/misc.xml new file mode 100644 index 0000000..db8786c --- /dev/null +++ b/.idea/misc.xml @@ -0,0 +1,7 @@ + + + + + + \ No newline at end of file diff --git a/.idea/modules.xml b/.idea/modules.xml new file mode 100644 index 0000000..8c614fe --- /dev/null +++ b/.idea/modules.xml @@ -0,0 +1,8 @@ + + + + + + + + \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..02ed4b2 --- /dev/null +++ b/__init__.py @@ -0,0 +1,12 @@ +from .node import * +from .install import * + +# A dictionary that contains all nodes you want to export with their names +# NOTE: names should be globally unique +NODE_CLASS_MAPPINGS = { + "AutomaticMask(segment anything)": AutomaticMask +} + +__all__ = ['NODE_CLASS_MAPPINGS'] + + diff --git a/install.py b/install.py new file mode 100644 index 0000000..2c14be2 --- /dev/null +++ b/install.py @@ -0,0 +1,17 @@ +import sys +import os.path +import subprocess + +custom_nodes_path = os.path.dirname(os.path.abspath(__file__)) + +def build_pip_install_cmds(args): + if "python_embeded" in sys.executable or "python_embedded" in sys.executable: + return [sys.executable, '-s', '-m', 'pip', 'install'] + args + else: + return [sys.executable, '-m', 'pip', 'install'] + args + +def ensure_package(): + cmds = build_pip_install_cmds(['-r', 'requirements.txt']) + subprocess.run(cmds, cwd=custom_nodes_path) + +ensure_package() \ No newline at end of file diff --git a/node.py b/node.py new file mode 100644 index 0000000..2002bab --- /dev/null +++ b/node.py @@ -0,0 +1,76 @@ +import sys +import os +import numpy as np +from PIL import Image +import torch +import matplotlib.pyplot as plt +import cv2 +from segment_anything import sam_model_registry, SamAutomaticMaskGenerator, SamPredictor +import folder_paths + +sys.path.append( + os.path.dirname(os.path.abspath(__file__)) +) + +def show_anns(anns, image_shape): + if len(anns) == 0: + return + sorted_anns = sorted(anns, key=(lambda x: x['area']), reverse=True) + + img = np.ones((image_shape[0], image_shape[1], 4)) + img[:,:,3] = 0 + for ann in sorted_anns: + m = ann['segmentation'] + color_mask = np.concatenate([np.random.random(3), [0.35]]) + img[m] = color_mask + # 将带有标注的numpy图像转化为torch张量 + annotated_img_tensor = torch.from_numpy(img) + + return annotated_img_tensor + +class AutomaticMask: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + }, + } + + # RETURN_NAMES = ("IMAGE",) + FUNCTION = "main" + CATEGORY = "segment_anything" + RETURN_TYPES = ("IMAGE",) + + def main(self, image): + sam_checkpoint = folder_paths.get_full_path('sams', 'sam_vit_h_4b8939.pth') + model_type = "vit_h" + + device = "cuda" + + sam = sam_model_registry[model_type](checkpoint=sam_checkpoint) + + sam.to(device=device) + + mask_generator = SamAutomaticMaskGenerator(sam) + + image_res = [] + for item in image: + image_shape = (item.shape[0], item.shape[1]) + print(image_shape) + item = Image.fromarray( + np.clip(255. * item.cpu().numpy(), 0, 255).astype(np.uint8)).convert('RGBA') + image_np = np.array(item) + image_np_rgb = image_np[..., :3] + + + # 生成蒙版 + masks = mask_generator.generate(image_np_rgb) + annotated_image_tensor = show_anns(masks, image_shape) + + image_res.append(annotated_image_tensor) + + return (image_res,) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..7688444 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,5 @@ +segment_anything +cv2import +matplotlib +torch +numpy \ No newline at end of file