添加根据点切割物体
This commit is contained in:
+111
-2
@@ -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"
|
||||
}
|
||||
|
||||
+21
-1
@@ -1,2 +1,22 @@
|
||||
open_clip_torch
|
||||
sentence_transformers
|
||||
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
|
||||
@@ -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]
|
||||
Reference in New Issue
Block a user