add mask from segmentation
This commit is contained in:
@@ -603,6 +603,58 @@ class MaskFromColor:
|
||||
|
||||
return (mask, )
|
||||
|
||||
class MaskFromSegmentation:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", ),
|
||||
"segments": ("INT", { "default": 6, "min": 1, "max": 16, "step": 1, }),
|
||||
"remove_isolated_pixels": ("INT", { "default": 0, "min": 0, "max": 32, "step": 1, }),
|
||||
"remove_small_masks": ("FLOAT", { "default": 0.0, "min": 0., "max": 1., "step": 0.01, }),
|
||||
"fill_holes": ("BOOLEAN", { "default": False }),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = "essentials"
|
||||
|
||||
def execute(self, image, segments, remove_isolated_pixels, fill_holes, remove_small_masks):
|
||||
im = image[0] # we only work on the first image in the batch
|
||||
im = Image.fromarray((im * 255).to(torch.uint8).cpu().numpy(), mode="RGB")
|
||||
im = im.quantize(palette=im.quantize(colors=segments), dither=Image.Dither.NONE)
|
||||
im = torch.tensor(np.array(im.convert("RGB"))).float() / 255.0
|
||||
|
||||
colors = im.reshape(-1, im.shape[-1])
|
||||
colors = torch.unique(colors, dim=0)
|
||||
|
||||
masks = []
|
||||
for color in colors:
|
||||
mask = (im == color).all(dim=-1).float()
|
||||
# remove isolated pixels
|
||||
if remove_isolated_pixels > 0:
|
||||
mask_np = mask.cpu().numpy()
|
||||
mask_np = scipy.ndimage.binary_opening(mask_np, structure=np.ones((remove_isolated_pixels, remove_isolated_pixels)))
|
||||
mask = torch.from_numpy(mask_np)
|
||||
|
||||
# fill holes
|
||||
if fill_holes:
|
||||
mask_np = mask.cpu().numpy()
|
||||
mask_np = scipy.ndimage.binary_fill_holes(mask_np)
|
||||
mask = torch.from_numpy(mask_np)
|
||||
|
||||
# if the mask is too small, it's probably noise
|
||||
if mask.sum() / (mask.shape[0]*mask.shape[1]) > remove_small_masks:
|
||||
masks.append(mask)
|
||||
|
||||
if masks == []:
|
||||
masks.append(torch.zeros_like(im).squeeze(-1).unsqueeze(0)) # return an empty mask if no masks were found, prevents errors
|
||||
|
||||
mask = torch.stack(masks, dim=0).float()
|
||||
|
||||
return (mask, )
|
||||
|
||||
class MaskFromBatch:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -1687,6 +1739,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"MaskFromColor+": MaskFromColor,
|
||||
"MaskFromBatch+": MaskFromBatch,
|
||||
"MaskBoundingBox+": MaskBoundingBox,
|
||||
"MaskFromSegmentation+": MaskFromSegmentation,
|
||||
|
||||
"SimpleMath+": SimpleMath,
|
||||
"ConsoleDebug+": ConsoleDebug,
|
||||
@@ -1737,6 +1790,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MaskFromColor+": "🔧 Mask From Color",
|
||||
"MaskFromBatch+": "🔧 Mask From Batch",
|
||||
"MaskBoundingBox+": "🔧 Mask Bounding Box",
|
||||
"MaskFromSegmentation+": "🔧 Mask From Segmentation",
|
||||
|
||||
"SimpleMath+": "🔧 Simple Math",
|
||||
"ConsoleDebug+": "🔧 Console Debug",
|
||||
|
||||
Reference in New Issue
Block a user