diff --git a/essentials.py b/essentials.py index 22f5c14..4e4372a 100644 --- a/essentials.py +++ b/essentials.py @@ -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",