From dd351769a032eceef69fe3cd6ff609264c3d1296 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E9=9B=AA=E5=B3=B0?= Date: Thu, 21 Dec 2023 19:16:46 +0800 Subject: [PATCH] add SamAutoMaskSEGS --- easyapi/SamNode.py | 54 ++++++++++++++++++++++++++++++++++++++++++++++ easyapi/util.py | 13 +++++++++++ requirements.txt | 1 + 3 files changed, 68 insertions(+) create mode 100644 easyapi/SamNode.py create mode 100644 easyapi/util.py create mode 100644 requirements.txt diff --git a/easyapi/SamNode.py b/easyapi/SamNode.py new file mode 100644 index 0000000..57ff93f --- /dev/null +++ b/easyapi/SamNode.py @@ -0,0 +1,54 @@ +from segment_anything import SamAutomaticMaskGenerator +import json +import numpy as np + +from util import tensor_to_pil + +class SamAutoMaskSEGS: + @classmethod + def INPUT_TYPES(self): + return {"required": { + "sam_model": ('SAM_MODEL', {}), + "image": ('IMAGE', {}), + "output_mode": (['uncompressed_rle', 'coco_rel'], {"default": "uncompressed_rle"}), + }, + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("RLE_SEGS", ) + + FUNCTION = "generate" + + OUTPUT_NODE = True + CATEGORY = "EasyApi/Detect" + + # INPUT_IS_LIST = False + # OUTPUT_IS_LIST = (False, False) + + def generate(self, sam_model, image, output_mode): + # 判断是不是HQ + encodeClassName = sam_model.image_encoder.__class__.__name__ + if encodeClassName == "ImageEncoderViTHQ": + from custom_nodes.comfyui_segment_anything.sam_hq.automatic import SamAutomaticMaskGeneratorHQ + from custom_nodes.comfyui_segment_anything.sam_hq.predictor import SamPredictorHQ + samHQ = SamPredictorHQ(sam_model, True) + mask_generator = SamAutomaticMaskGeneratorHQ(samHQ, output_mode=output_mode) + else: + mask_generator = SamAutomaticMaskGenerator(sam_model, output_mode=output_mode) + image_pil = tensor_to_pil(image) + image_np = np.array(image_pil) + image_np_rgb = image_np[..., :3] + + masks = mask_generator.generate(image_np_rgb) + masksRle = json.JSONEncoder().encode(masks) + return {"ui": {"segsRle": (masksRle,)}, "result": (masksRle,)} + + +NODE_CLASS_MAPPINGS = { + "SamAutoMaskSEGS": SamAutoMaskSEGS, +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "SamAutoMaskSEGS": "SamAutoMaskSEGS", +} diff --git a/easyapi/util.py b/easyapi/util.py new file mode 100644 index 0000000..af91268 --- /dev/null +++ b/easyapi/util.py @@ -0,0 +1,13 @@ +import numpy as np +import torch +from PIL import Image + +# Tensor to PIL +def tensor_to_pil(image): + return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8)) + + +# Convert PIL to Tensor +def pil_2_tensor(image): + return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) + diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..4032d22 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +segment_anything \ No newline at end of file