From de54851c0f4a58b86892ca656780416233b1396f Mon Sep 17 00:00:00 2001 From: AbyssYuan0 Date: Fri, 5 Jan 2024 17:06:34 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E6=A0=B9=E6=8D=AE=E7=82=B9?= =?UTF-8?q?=E5=88=87=E5=89=B2=E7=89=A9=E4=BD=93?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- __init__.py | 113 ++++++++++++++++++++++++++++++++++++++++++++++- requirements.txt | 22 ++++++++- seg.py | 45 +++++++++++++++++++ 3 files changed, 177 insertions(+), 3 deletions(-) create mode 100644 seg.py diff --git a/__init__.py b/__init__.py index 05593cc..90e90f6 100644 --- a/__init__.py +++ b/__init__.py @@ -6,6 +6,7 @@ import numpy as np import torch import comfy.utils from .videoCut import getCutList, saveToDir +from .seg import get_mask def getImageSize(IMAGE) -> tuple[int, int]: @@ -27,6 +28,24 @@ def imgToTensor(img): return imaget +def img_to_mask(mask): + mask = mask.convert("RGBA") + mask = np.array(mask.getchannel('R')).astype(np.float32) / 255.0 + mask = torch.from_numpy(mask) + return mask + + +def img_to_np(img): + if img.mode == "RGBA": + img = img.convert("RGB") + img = np.array(img) + return img + + +def np_to_img(numpy): + return Image.fromarray(numpy.astype(np.uint8)) + + class ImageOverlap: def __init__(self): @@ -499,6 +518,95 @@ class mkdir: return (new_dir_path,) +class findCenterOfMask: + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "mask": ("MASK",), + } + } + + CATEGORY = "badger" + + RETURN_TYPES = ("FLOAT", "FLOAT",) + RETURN_NAMES = ("X", "Y",) + FUNCTION = "find_center_of_mask" + + def find_center_of_mask(self, mask): + if mask.dim() == 3: + mask = mask.squeeze(0) # Remove the channel dimension if it exists + assert mask.dim() == 2, "Mask must be 2D" + + # Create grids for x and y coordinates + h, w = mask.size() + x_coords = torch.arange(w).float().to(mask.device) + y_coords = torch.arange(h).float().to(mask.device) + + # Compute the center of mass (centroid) of the mask + total_mass = mask.sum() + if total_mass > 0: + x_center = (mask.sum(dim=0) * x_coords).sum() / total_mass + y_center = (mask.sum(dim=1) * y_coords).sum() / total_mass + else: + x_center, y_center = torch.tensor(0), torch.tensor(0) + + # Convert to int + X = float(x_center.item()) + Y = float(y_center.item()) + + return (X, Y,) + + +class SegmentToMaskByPoint: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "img": ("IMAGE",), + "X": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 4096.0, + "step": 0.1, + "display": "number" + }), + "Y": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 4096.0, + "step": 0.1, + "display": "number" + }), + "dilate": ("INT", { + "default": 15, + "min": 0, + "max": 4096.0, + "step": 1, + "display": "number" + }), + "sam_ckpt": ("SAM_MODEL",), + } + } + + CATEGORY = "badger" + + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("mask",) + FUNCTION = "seg_to_mask_by_point" + + def seg_to_mask_by_point(self, img, X, Y, dilate, sam_ckpt): + img = tensorToImg(img) + img = img_to_np(img) + latest_coords = [X, Y] + mask = get_mask(img, latest_coords, dilate, sam_ckpt) + mask = np_to_img(mask) + mask = img_to_mask(mask) + + return (mask.unsqueeze(0),) + + NODE_CLASS_MAPPINGS = { "ImageOverlap-badger": ImageOverlap, "FloatToInt-badger": FloatToInt, @@ -511,9 +619,10 @@ NODE_CLASS_MAPPINGS = { "getImageSide-badger": getImageSide, "VideoCut-badger": videoCut, "getParentDir-badger": getParentDir, - "mkdir-badger": mkdir + "mkdir-badger": mkdir, + "findCenterOfMask-badger": findCenterOfMask, + "SegmentToMaskByPoint-badger": SegmentToMaskByPoint, } NODE_DISPLAY_NAME_MAPPINGS = { - "ImageOverlap": "Example test" } diff --git a/requirements.txt b/requirements.txt index 20f4431..e97670d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,2 +1,22 @@ open_clip_torch -sentence_transformers \ No newline at end of file +sentence_transformers +pyyaml +tqdm +numpy +easydict==1.9.0 +scikit-image +scikit-learn +opencv-python +tensorflow +joblib +matplotlib +pandas +albumentations==0.5.2 +hydra-core==1.1.0 +pytorch-lightning==1.2.9 +tabulate +kornia==0.5.0 +webdataset +packaging +wldhx.yadisk-direct +timm \ No newline at end of file diff --git a/seg.py b/seg.py new file mode 100644 index 0000000..a40d7cf --- /dev/null +++ b/seg.py @@ -0,0 +1,45 @@ +import numpy as np +import cv2 +from segment_anything import SamPredictor + + +def dilate_mask(mask, dilate_factor=15): + mask = mask.astype(np.uint8) + mask = cv2.dilate( + mask, + np.ones((dilate_factor, dilate_factor), np.uint8), + iterations=1 + ) + return mask + + +def predict_masks_with_sam( + img, + point_coords, + point_labels, + ckpt_p, +): + point_coords = np.array(point_coords) + point_labels = np.array(point_labels) + predictor = SamPredictor(ckpt_p) + + predictor.set_image(img) + masks, scores, logits = predictor.predict( + point_coords=point_coords, + point_labels=point_labels, + multimask_output=True, + ) + return masks, scores, logits + + +def get_mask(img, latest_coords, dilate_kernel_size, sam_ckpt): + masks, _, _ = predict_masks_with_sam( + img, + [latest_coords], + [1], + sam_ckpt, + ) + masks = masks.astype(np.uint8) * 255 + + masks = [dilate_mask(mask, dilate_kernel_size) for mask in masks] + return masks[1]