From 074469237f1becb7d79296126e428c3d780158d8 Mon Sep 17 00:00:00 2001 From: Waheed-Mousad Date: Tue, 7 Jul 2026 16:43:01 +0300 Subject: [PATCH] working nodes --- .gitignore | 4 + README.md | 0 __init__.py | 35 +++++ nodes/__init__.py | 0 nodes/convert.py | 74 ++++++++++ nodes/dedup.py | 104 +++++++++++++ nodes/expand.py | 152 +++++++++++++++++++ nodes/pick.py | 56 +++++++ nodes/sort.py | 86 +++++++++++ utils/__init__.py | 0 utils/mask_ops.py | 368 ++++++++++++++++++++++++++++++++++++++++++++++ 11 files changed, 879 insertions(+) create mode 100644 .gitignore create mode 100644 README.md create mode 100644 __init__.py create mode 100644 nodes/__init__.py create mode 100644 nodes/convert.py create mode 100644 nodes/dedup.py create mode 100644 nodes/expand.py create mode 100644 nodes/pick.py create mode 100644 nodes/sort.py create mode 100644 utils/__init__.py create mode 100644 utils/mask_ops.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..52c37f4 --- /dev/null +++ b/.gitignore @@ -0,0 +1,4 @@ +/.idea +/__pycache__ +/nodes/__pycache__ +/utils/__pycache__ diff --git a/README.md b/README.md new file mode 100644 index 0000000..e69de29 diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..6059e39 --- /dev/null +++ b/__init__.py @@ -0,0 +1,35 @@ +from .nodes.convert import ( + NODE_CLASS_MAPPINGS as convert_class_mappings, + NODE_DISPLAY_NAME_MAPPINGS as convert_display_mappings, +) +from .nodes.sort import ( + NODE_CLASS_MAPPINGS as sort_class_mappings, + NODE_DISPLAY_NAME_MAPPINGS as sort_display_mappings, +) +from .nodes.pick import ( + NODE_CLASS_MAPPINGS as pick_class_mappings, + NODE_DISPLAY_NAME_MAPPINGS as pick_display_mappings, +) +from .nodes.dedup import ( + NODE_CLASS_MAPPINGS as dedup_class_mappings, + NODE_DISPLAY_NAME_MAPPINGS as dedup_display_mappings, +) +from .nodes.expand import ( + NODE_CLASS_MAPPINGS as expand_class_mappings, + NODE_DISPLAY_NAME_MAPPINGS as expand_display_mappings, +) + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +for cm, dm in [ + (convert_class_mappings, convert_display_mappings), + (sort_class_mappings, sort_display_mappings), + (pick_class_mappings, pick_display_mappings), + (dedup_class_mappings, dedup_display_mappings), + (expand_class_mappings, expand_display_mappings), +]: + NODE_CLASS_MAPPINGS.update(cm) + NODE_DISPLAY_NAME_MAPPINGS.update(dm) + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/nodes/__init__.py b/nodes/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/nodes/convert.py b/nodes/convert.py new file mode 100644 index 0000000..935f660 --- /dev/null +++ b/nodes/convert.py @@ -0,0 +1,74 @@ +import torch + + +class MaskBatchToList: + """Convert a mask batch [N, H, W] into a list of N individual [1, H, W] masks.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "masks": ("MASK",), + } + } + + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("masks",) + OUTPUT_IS_LIST = (True,) + FUNCTION = "batch_to_list" + CATEGORY = "MultiMaskOps" + + def batch_to_list(self, masks): + if masks.dim() == 2: + masks = masks.unsqueeze(0) # [H, W] -> [1, H, W] + + mask_list = [masks[i:i + 1] for i in range(masks.shape[0])] + return (mask_list,) + + +class MaskListToBatch: + """Convert a list of masks into a single batch tensor [N, H, W].""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "masks": ("MASK",), + } + } + + INPUT_IS_LIST = True + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("masks",) + FUNCTION = "list_to_batch" + CATEGORY = "MultiMaskOps" + + def list_to_batch(self, masks): + normalized = [] + for m in masks: + if m.dim() == 2: + m = m.unsqueeze(0) # [H, W] -> [1, H, W] + normalized.append(m) + + target_h, target_w = normalized[0].shape[-2:] + resized = [] + for m in normalized: + if m.shape[-2:] != (target_h, target_w): + m = torch.nn.functional.interpolate( + m.unsqueeze(1), size=(target_h, target_w), mode="nearest" + ).squeeze(1) + resized.append(m) + + batch = torch.cat(resized, dim=0) # -> [N, H, W] + return (batch,) + + +NODE_CLASS_MAPPINGS = { + "MaskBatchToList": MaskBatchToList, + "MaskListToBatch": MaskListToBatch, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "MaskBatchToList": "Mask Batch To List", + "MaskListToBatch": "Mask List To Batch", +} \ No newline at end of file diff --git a/nodes/dedup.py b/nodes/dedup.py new file mode 100644 index 0000000..88db0fe --- /dev/null +++ b/nodes/dedup.py @@ -0,0 +1,104 @@ +import torch + +from ..utils.mask_ops import ( + mask_to_bhw, + group_by_iou, + mask_union, + mask_intersection, + colored_group_image, + distinct_colors, +) + + +class MaskDeduplicate: + """Detect and remove duplicate masks from a list using transitive IoU + grouping. Each group of near-identical masks is collapsed into a single + mask according to the chosen keep mode.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "masks": ("MASK",), + "iou_threshold": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01}), + "keep": (["largest", "smallest", "first", "AND", "OR"],), + "threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + } + } + + INPUT_IS_LIST = True + RETURN_TYPES = ("MASK", "INT", "IMAGE") + RETURN_NAMES = ("masks", "removed_count", "group_images") + OUTPUT_IS_LIST = (True, False, True) + FUNCTION = "deduplicate" + CATEGORY = "MultiMaskOps" + + def deduplicate(self, masks, iou_threshold, keep, threshold): + # INPUT_IS_LIST: collapse widget lists to scalars. + iou_threshold = iou_threshold[0] + keep = keep[0] + threshold = threshold[0] + + # Flatten to individual [1, H, W] masks (originals preserved). + flat = [] + for m in masks: + m = mask_to_bhw(m) + for j in range(m.shape[0]): + flat.append(m[j:j + 1]) + + if len(flat) == 0: + raise ValueError("MaskDeduplicate received an empty mask list.") + + batch = torch.cat(flat, dim=0) # [N, H, W] + n = batch.shape[0] + + # Transitive IoU grouping (connected components). + groups = group_by_iou(batch, iou_threshold, binarize_threshold=threshold) + + result_masks = [] + group_images = [] + + for group in groups: + group_masks = [batch[i] for i in group] # list of [H, W] + + # Collapse the group to a single survivor per keep mode. + if keep == "first": + survivor = batch[group[0]] + elif keep == "largest": + areas = [(batch[i] > threshold).sum().item() for i in group] + survivor = batch[group[int(torch.tensor(areas).argmax())]] + elif keep == "smallest": + areas = [(batch[i] > threshold).sum().item() for i in group] + survivor = batch[group[int(torch.tensor(areas).argmin())]] + elif keep == "OR": + survivor = mask_union(group_masks) + elif keep == "AND": + survivor = mask_intersection(group_masks) + else: + survivor = batch[group[0]] + + result_masks.append(survivor.unsqueeze(0)) # [1, H, W] + + # Debug image only for groups that actually had duplicates (2+). + if len(group) >= 2: + colors = distinct_colors(len(group)) + group_images.append(colored_group_image(group_masks, colors)) + + removed_count = n - len(groups) + + # If no duplicate groups existed, emit a single blank image so the + # IMAGE output is never an empty list (which some nodes dislike). + if len(group_images) == 0: + H, W = batch.shape[-2:] + group_images.append(torch.ones(1, H, W, 3)) + + return (result_masks, removed_count, group_images) + + +NODE_CLASS_MAPPINGS = { + "MaskDeduplicate": MaskDeduplicate, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "MaskDeduplicate": "Mask Deduplicate", +} \ No newline at end of file diff --git a/nodes/expand.py b/nodes/expand.py new file mode 100644 index 0000000..c7eac8e --- /dev/null +++ b/nodes/expand.py @@ -0,0 +1,152 @@ +import torch + +from ..utils.mask_ops import ( + mask_to_bhw, + binarize, + dilate_batch, + feather, + subtract, +) + + +class MaskExpandWithoutOverlap: + """Expand (and optionally feather) each mask in a list while preventing it + from leaking into other masks' regions. + + overlap_policy controls what is protected and how contested empty space is + handled: + - allow : subtract only others' ORIGINAL regions; expansions may + overlap freely in empty space (values kept). + - retreat : subtract others' EXPANDED regions; contested pixels cleared + in both masks (clean empty seam). + - priority : like retreat, but earlier-in-list masks win contested + pixels (gap-free, deterministic by order). + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "masks": ("MASK",), + "expand_pixels": ("INT", {"default": 10, "min": 0, "max": 4096, "step": 1}), + "tapered_corners": ("BOOLEAN", {"default": True}), + "feather_amount": ("INT", {"default": 0, "min": 0, "max": 4096, "step": 1}), + "feather_mode": (["cut_feather", "preserve_feather"],), + "overlap_policy": (["allow", "retreat", "priority"],), + "subtract_hard": ("BOOLEAN", {"default": True}), + "threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + } + } + + INPUT_IS_LIST = True + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("masks",) + OUTPUT_IS_LIST = (True,) + FUNCTION = "expand" + CATEGORY = "MultiMaskOps" + + def expand(self, masks, expand_pixels, tapered_corners, feather_amount, feather_mode, + overlap_policy, subtract_hard, threshold): + expand_pixels = expand_pixels[0] + tapered_corners = tapered_corners[0] + feather_amount = feather_amount[0] + feather_mode = feather_mode[0] + overlap_policy = overlap_policy[0] + subtract_hard = subtract_hard[0] + threshold = threshold[0] + + # Flatten to a single [N, H, W] batch. + flat = [] + for m in masks: + m = mask_to_bhw(m) + for j in range(m.shape[0]): + flat.append(m[j]) + n = len(flat) + if n == 0: + raise ValueError("MaskExpandWithoutOverlap received an empty mask list.") + + batch = torch.stack(flat, dim=0) # [N, H, W] originals + + # --- Precompute shared quantities ONCE (this is the speed win) --- + # Hard originals and their cumulative union. + orig_hard = binarize(batch, threshold) if subtract_hard else batch # [N,H,W] + total_orig_union = orig_hard.max(dim=0).values # [H, W] union of all originals + + # Expanded masks: one batched separable dilation for the whole stack. + expanded = dilate_batch(batch, expand_pixels, tapered_corners) # [N, H, W] soft + expanded_hard = binarize(expanded, threshold) # [N, H, W] + total_exp_union = expanded_hard.max(dim=0).values # [H, W] + + results = [] + for i in range(n): + self_orig_hard = orig_hard[i] + + # "others original" = union of all originals minus this one's. + # Since union is max, removing self where self is the only contributor + # requires care; simplest correct form: max over others. + if n == 1: + others_original = torch.zeros_like(self_orig_hard) + others_expanded = torch.zeros_like(self_orig_hard) + else: + others_original = _union_excluding(orig_hard, i) + if overlap_policy == "allow": + others_expanded = others_original # unused, kept defined + elif overlap_policy == "retreat": + others_expanded = _union_excluding(expanded_hard, i) + else: # priority: only earlier masks protect their expansion + if i == 0: + others_expanded = torch.zeros_like(self_orig_hard) + else: + earlier_exp = expanded_hard[:i].max(dim=0).values + earlier_orig = orig_hard[:i].max(dim=0).values + others_expanded = torch.maximum(earlier_exp, earlier_orig) + + # Region to subtract during processing. + if overlap_policy == "allow": + region = others_original + else: + region = others_expanded + + grown = expanded[i] # already dilated + + if feather_amount > 0: + if feather_mode == "cut_feather": + soft = feather(grown, feather_amount) + final = subtract(soft, region) + else: # preserve_feather + cut = subtract(grown, region) + soft = feather(cut, feather_amount) + final = subtract(soft, region) + else: + final = subtract(grown, region) + + # INVARIANT: never leak into any other mask's ORIGINAL area. + if n > 1: + final = subtract(final, others_original) + + results.append(final.unsqueeze(0)) + + return (results,) + + +def _union_excluding(stack_nhw, i): + """Max-union of all slices except index i, without rebuilding lists.""" + n = stack_nhw.shape[0] + if n == 1: + return torch.zeros_like(stack_nhw[0]) + if i == 0: + return stack_nhw[1:].max(dim=0).values + if i == n - 1: + return stack_nhw[:-1].max(dim=0).values + top = stack_nhw[:i].max(dim=0).values + bot = stack_nhw[i + 1:].max(dim=0).values + return torch.maximum(top, bot) + + +NODE_CLASS_MAPPINGS = { + "MaskExpandWithoutOverlap": MaskExpandWithoutOverlap, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "MaskExpandWithoutOverlap": "Mask Expand Without Overlap", +} \ No newline at end of file diff --git a/nodes/pick.py b/nodes/pick.py new file mode 100644 index 0000000..dce0043 --- /dev/null +++ b/nodes/pick.py @@ -0,0 +1,56 @@ +from ..utils.mask_ops import mask_to_bhw + + +class MaskPickByIndex: + """Return a single mask from a list by index. + + Supports Python-style negative indexing (-1 = last mask, -2 = second to + last, ...). Out-of-range indices are clamped to the nearest valid mask on + both ends, so the node always returns a valid mask given a non-empty list. + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "masks": ("MASK",), + "index": ("INT", {"default": 0, "min": -10000, "max": 10000, "step": 1}), + } + } + + INPUT_IS_LIST = True + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("mask",) + FUNCTION = "pick" + CATEGORY = "MultiMaskOps" + + def pick(self, masks, index): + # INPUT_IS_LIST: 'index' arrives as a single-item list. + index = index[0] + + # Flatten any batched entries into individual [1, H, W] masks. + flat = [] + for m in masks: + m = mask_to_bhw(m) + for j in range(m.shape[0]): + flat.append(m[j:j + 1]) + + n = len(flat) + if n == 0: + raise ValueError("MaskPickByIndex received an empty mask list.") + + # Negative index counts from the end. + resolved = index + n if index < 0 else index + # Clamp to a valid position on both ends. + resolved = max(0, min(resolved, n - 1)) + + return (flat[resolved],) + + +NODE_CLASS_MAPPINGS = { + "MaskPickByIndex": MaskPickByIndex, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "MaskPickByIndex": "Mask Pick By Index", +} \ No newline at end of file diff --git a/nodes/sort.py b/nodes/sort.py new file mode 100644 index 0000000..f004a98 --- /dev/null +++ b/nodes/sort.py @@ -0,0 +1,86 @@ +import torch + +from ..utils.mask_ops import ( + mask_to_bhw, + reference_points, + sort_indices, + debug_image, +) + + +class MaskSortByPosition: + """Sort a list of masks by spatial position (centroid or extreme point).""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "masks": ("MASK",), + "reference_point": (["centroid", "extreme_point"],), + "sort_left_to_right": ("BOOLEAN", {"default": True}), + "sort_top_to_bottom": ("BOOLEAN", {"default": True}), + "primary_axis_horizontal": ("BOOLEAN", {"default": True}), + "flatten_before_sort": ("BOOLEAN", {"default": False}), + "drop_below_threshold": ("BOOLEAN", {"default": False}), + "threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + } + } + + INPUT_IS_LIST = True + RETURN_TYPES = ("MASK", "IMAGE") + RETURN_NAMES = ("masks", "debug_image") + OUTPUT_IS_LIST = (True, False) + FUNCTION = "sort_masks" + CATEGORY = "MultiMaskOps" + + def sort_masks(self, masks, reference_point, sort_left_to_right, + sort_top_to_bottom, primary_axis_horizontal, + flatten_before_sort, drop_below_threshold, threshold): + # INPUT_IS_LIST: widgets arrive as single-item lists; collapse them. + reference_point = reference_point[0] + sort_left_to_right = sort_left_to_right[0] + sort_top_to_bottom = sort_top_to_bottom[0] + primary_axis_horizontal = primary_axis_horizontal[0] + flatten_before_sort = flatten_before_sort[0] + drop_below_threshold = drop_below_threshold[0] + threshold = threshold[0] + + # Normalize incoming masks to individual [1, H, W], preserving originals. + flat = [] + for m in masks: + m = mask_to_bhw(m) + for j in range(m.shape[0]): + flat.append(m[j:j + 1]) + batch = torch.cat(flat, dim=0) # [N, H, W] - ORIGINAL, unmodified + + H, W = batch.shape[-2:] + + # Preprocessing happens inside reference_points and affects ONLY sorting. + pts = reference_points( + batch, primary_axis_horizontal, + sort_left_to_right, sort_top_to_bottom, reference_point, + threshold=threshold, + flatten_before_sort=flatten_before_sort, + drop_below_threshold=drop_below_threshold, + ) + order = sort_indices( + pts, primary_axis_horizontal, + sort_left_to_right, sort_top_to_bottom + ) + + # Return ORIGINAL masks, just reordered. + sorted_masks = [batch[i:i + 1] for i in order] + sorted_pts = [pts[i] for i in order] + + dbg = debug_image(sorted_pts, H, W, radius=10) + + return (sorted_masks, dbg) + + +NODE_CLASS_MAPPINGS = { + "MaskSortByPosition": MaskSortByPosition, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "MaskSortByPosition": "Mask Sort By Position", +} \ No newline at end of file diff --git a/utils/__init__.py b/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/utils/mask_ops.py b/utils/mask_ops.py new file mode 100644 index 0000000..0d9a73f --- /dev/null +++ b/utils/mask_ops.py @@ -0,0 +1,368 @@ +import torch + + +def mask_to_bhw(m): + """Normalize any MASK tensor to [B, H, W].""" + if m.dim() == 2: + return m.unsqueeze(0) + return m + + +def preprocess_for_sort(mask_hw, threshold, flatten_before_sort, drop_below_threshold): + """ + Produce a temporary mask used ONLY to determine the sorting reference point. + The caller must keep the original mask for output. + + Rules: + - flatten_before_sort=True: + binary mask: (pixel >= threshold) -> 1, else 0. + (drop_below_threshold is ignored in this case.) + - flatten_before_sort=False and drop_below_threshold=True: + thresholded ReLU: pixels < threshold -> 0, pixels >= threshold keep value. + - both False: + mask returned unchanged. + """ + if flatten_before_sort: + return (mask_hw >= threshold).float() + if drop_below_threshold: + return torch.where(mask_hw >= threshold, mask_hw, torch.zeros_like(mask_hw)) + return mask_hw + + +def weighted_centroid(mask_hw): + """ + Value-weighted centroid of a mask. Each pixel contributes proportionally + to its value, so soft/feathered masks are handled correctly and binary + masks reduce to the plain mean. + + Returns (x, y) or None if the mask has no positive weight. + """ + pos = mask_hw > 0 + if not pos.any(): + return None + ys, xs = torch.where(pos) + w = mask_hw[ys, xs].float() + total = w.sum() + if total <= 0: + return None + x = (xs.float() * w).sum() / total + y = (ys.float() * w).sum() / total + return (float(x), float(y)) + + +def extreme_point(mask_hw, primary_axis_horizontal, + sort_left_to_right, sort_top_to_bottom): + """ + Outermost valid pixel along the primary sorting axis/direction. + 'Valid' = any pixel > 0 in the (already preprocessed) mask. + + Returns (x, y) or None if no valid pixels. + """ + pos = mask_hw > 0 + if not pos.any(): + return None + ys, xs = torch.where(pos) + if primary_axis_horizontal: + idx = xs.argmin() if sort_left_to_right else xs.argmax() + else: + idx = ys.argmin() if sort_top_to_bottom else ys.argmax() + return (float(xs[idx].item()), float(ys[idx].item())) + + +def reference_points(masks_bhw, primary_axis_horizontal, + sort_left_to_right, sort_top_to_bottom, mode, + threshold=0.5, flatten_before_sort=False, + drop_below_threshold=False): + """ + Compute one (x, y) reference point per mask, applying preprocessing that + affects ONLY the reference-point calculation (never the output masks). + """ + pts = [] + for i in range(masks_bhw.shape[0]): + proc = preprocess_for_sort( + masks_bhw[i], threshold, flatten_before_sort, drop_below_threshold + ) + if mode == "centroid": + p = weighted_centroid(proc) + else: # extreme_point + p = extreme_point( + proc, primary_axis_horizontal, + sort_left_to_right, sort_top_to_bottom + ) + if p is None: + p = (0.0, 0.0) # empty mask sorts predictably at origin + pts.append(p) + return pts + + +def sort_indices(points, primary_axis_horizontal, + sort_left_to_right, sort_top_to_bottom): + """Stable sorted order of indices given reference points.""" + n = len(points) + + def x_key(x): + return x if sort_left_to_right else -x + + def y_key(y): + return y if sort_top_to_bottom else -y + + if primary_axis_horizontal: + key = lambda i: (x_key(points[i][0]), y_key(points[i][1]), i) + else: + key = lambda i: (y_key(points[i][1]), x_key(points[i][0]), i) + + return sorted(range(n), key=key) + + +def expand_points_mask(points, height, width, radius=10): + """[H, W] tensor with a filled square (~dot) of given radius at each point.""" + canvas = torch.zeros(height, width) + for (x, y) in points: + xi, yi = int(round(x)), int(round(y)) + x0 = max(0, xi - radius) + x1 = min(width, xi + radius + 1) + y0 = max(0, yi - radius) + y1 = min(height, yi + radius + 1) + if x1 > x0 and y1 > y0: + canvas[y0:y1, x0:x1] = 1.0 + return canvas + + +def debug_image(points, height, width, radius=10): + """White [1, H, W, 3] IMAGE with black dots at each reference point.""" + dots = expand_points_mask(points, height, width, radius) + img = torch.ones(height, width, 3) + img[dots > 0.5] = 0.0 + return img.unsqueeze(0) + + +# --------------------------------------------------------------------------- +# Shared mask logic (used by dedup, and by future expand/blur nodes) +# --------------------------------------------------------------------------- + +def binarize(mask_hw, threshold): + """(pixel >= threshold) -> 1.0, else 0.0.""" + return (mask_hw >= threshold).float() + + +def mask_union(masks): + """Soft OR across a list/stack of masks: elementwise max. + Reduces to boolean OR on binary masks, preserves soft edges otherwise.""" + stacked = torch.stack([m for m in masks], dim=0) + return stacked.max(dim=0).values + + +def mask_intersection(masks): + """Soft AND across a list/stack of masks: elementwise min. + Reduces to boolean AND on binary masks, preserves soft edges otherwise.""" + stacked = torch.stack([m for m in masks], dim=0) + return stacked.min(dim=0).values + + +def iou(mask_a_hw, mask_b_hw, threshold=0.5): + """Intersection-over-Union of two masks, computed on binarized versions. + Returns a float in [0, 1]. Two empty masks are treated as identical (1.0).""" + a = binarize(mask_a_hw, threshold) > 0 + b = binarize(mask_b_hw, threshold) > 0 + inter = (a & b).sum().item() + union = (a | b).sum().item() + if union == 0: + return 1.0 # both empty -> identical + return inter / union + + +def iou_matrix(masks_bhw, threshold=0.5): + """Pairwise IoU matrix [N, N] for a stack of masks.""" + n = masks_bhw.shape[0] + mat = torch.zeros(n, n) + for i in range(n): + mat[i, i] = 1.0 + for j in range(i + 1, n): + v = iou(masks_bhw[i], masks_bhw[j], threshold) + mat[i, j] = v + mat[j, i] = v + return mat + + +def connected_components(adjacency): + """Union-find connected components from a boolean [N, N] adjacency matrix. + Returns a list of groups, each a sorted list of indices. Group order + follows first appearance for determinism.""" + n = adjacency.shape[0] + parent = list(range(n)) + + def find(x): + while parent[x] != x: + parent[x] = parent[parent[x]] + x = parent[x] + return x + + def union(a, b): + ra, rb = find(a), find(b) + if ra != rb: + parent[max(ra, rb)] = min(ra, rb) + + for i in range(n): + for j in range(i + 1, n): + if adjacency[i][j]: + union(i, j) + + groups = {} + for i in range(n): + root = find(i) + groups.setdefault(root, []).append(i) + + # deterministic ordering: by smallest index in each group + return [sorted(g) for g in sorted(groups.values(), key=lambda g: min(g))] + + +def group_by_iou(masks_bhw, iou_threshold, binarize_threshold=0.5): + """Group masks by transitive IoU similarity (connected components). + Returns list of index-groups.""" + mat = iou_matrix(masks_bhw, binarize_threshold) + adj = mat >= iou_threshold + return connected_components(adj) + + +def distinct_colors(n): + """Generate n visually distinct RGB colors (0-1 floats) by spreading hues.""" + import colorsys + colors = [] + for i in range(max(1, n)): + h = (i / max(1, n)) % 1.0 + s = 0.65 + v = 0.95 + r, g, b = colorsys.hsv_to_rgb(h, s, v) + colors.append((r, g, b)) + return colors + + +def colored_group_image(group_masks, colors=None): + """Render one group's masks onto a white [1, H, W, 3] image, each mask in + its own color. Overlapping pixels are the value-weighted average of the + contributing colors composited over white, so overlaps read as a blend + (e.g. blue + red -> purple) rather than one color hiding the other. + + group_masks: list of [1, H, W] or [H, W] mask tensors (soft values ok). + """ + masks = [] + for m in group_masks: + if m.dim() == 3: + m = m[0] + masks.append(m) + H, W = masks[0].shape[-2:] + + if colors is None: + colors = distinct_colors(len(masks)) + + # Accumulate weighted color and total weight per pixel. + color_accum = torch.zeros(H, W, 3) + weight_accum = torch.zeros(H, W) + for m, col in zip(masks, colors): + w = m.clamp(0.0, 1.0) # [H, W] weight + col_t = torch.tensor(col, dtype=torch.float32) # [3] + color_accum += w.unsqueeze(-1) * col_t.view(1, 1, 3) + weight_accum += w + + # Mean color where covered; white elsewhere. + covered = weight_accum > 0 + mean_color = torch.ones(H, W, 3) + safe_w = weight_accum.clamp(min=1e-6).unsqueeze(-1) + mean_color_vals = color_accum / safe_w + # Composite over white by the max coverage so soft edges fade to white. + # alpha = how strongly any mask covers this pixel (cap at 1). + alpha = weight_accum.clamp(0.0, 1.0).unsqueeze(-1) + blended = mean_color_vals * alpha + torch.ones(H, W, 3) * (1 - alpha) + img = torch.where(covered.unsqueeze(-1), blended, mean_color) + return img.unsqueeze(0) + + +# --------------------------------------------------------------------------- +# Expansion / feathering (used by the expand-without-overlap node) +# --------------------------------------------------------------------------- + +def _grow_kernel(tapered_corners): + """3x3 structuring element matching ComfyUI's native GrowMask. + tapered_corners=True zeroes the diagonal corners for rounder growth.""" + import numpy as np + c = 0 if tapered_corners else 1 + return np.array([[c, 1, c], + [1, 1, 1], + [c, 1, c]]) + + +def dilate(mask_hw, pixels, tapered_corners=True): + """Grow (or shrink) a mask using ComfyUI's native method: iterative 3x3 + grey-dilation via scipy, repeated abs(pixels) times. Positive pixels + dilate, negative erode. tapered_corners rounds the growth. This matches + the native GrowMask node's output exactly. + + Falls back to a separable torch max-pool if scipy is unavailable.""" + if pixels == 0: + return mask_hw + try: + import numpy as np + import scipy.ndimage + kernel = _grow_kernel(tapered_corners) + out = mask_hw.detach().cpu().numpy() + for _ in range(abs(int(pixels))): + if pixels < 0: + out = scipy.ndimage.grey_erosion(out, footprint=kernel) + else: + out = scipy.ndimage.grey_dilation(out, footprint=kernel) + return torch.from_numpy(out).to(mask_hw.device, mask_hw.dtype) + except Exception: + # Fallback: separable square max-pool (fast, but square corners). + if pixels <= 0: + return mask_hw + k = 2 * int(pixels) + 1 + p = int(pixels) + x = mask_hw.unsqueeze(0).unsqueeze(0) + x = torch.nn.functional.max_pool2d(x, (1, k), stride=1, padding=(0, p)) + x = torch.nn.functional.max_pool2d(x, (k, 1), stride=1, padding=(p, 0)) + return x[0, 0] + + +def dilate_batch(masks_nhw, pixels, tapered_corners=True): + """Apply native-style dilation to each mask in an [N, H, W] stack.""" + if pixels == 0: + return masks_nhw + outs = [dilate(masks_nhw[i], pixels, tapered_corners) for i in range(masks_nhw.shape[0])] + return torch.stack(outs, dim=0) + + +def feather(mask_hw, radius): + """Soft-edge feather using the same method as ComfyUI/KJNodes: a + torchvision GaussianBlur where the blur value is used as sigma and the + kernel size is 6*sigma+1 (forced odd). radius<=0 returns the mask + unchanged. Matches GrowMaskWithBlur's blur_radius behavior. + """ + if radius <= 0: + return mask_hw + from torchvision.transforms import functional as TF + sigma = float(radius) + kernel_size = int(6 * int(sigma) + 1) + if kernel_size % 2 == 0: + kernel_size += 1 + x = mask_hw.unsqueeze(0) # [1, H, W] -> treated as (C,H,W) by TF + x = TF.gaussian_blur(x, kernel_size=[kernel_size, kernel_size], sigma=[sigma, sigma]) + return x[0].clamp(0.0, 1.0) + + +def feather_batch(masks_nhw, radius): + """Batched feather matching the native GaussianBlur convention.""" + if radius <= 0: + return masks_nhw + from torchvision.transforms import functional as TF + sigma = float(radius) + kernel_size = int(6 * int(sigma) + 1) + if kernel_size % 2 == 0: + kernel_size += 1 + x = masks_nhw.unsqueeze(1) # [N,1,H,W] + x = TF.gaussian_blur(x, kernel_size=[kernel_size, kernel_size], sigma=[sigma, sigma]) + return x[:, 0].clamp(0.0, 1.0) + + +def subtract(mask_hw, remove_hw): + """Subtract remove_hw from mask_hw: mask * (1 - remove), clamped to [0,1].""" + return (mask_hw * (1.0 - remove_hw.clamp(0.0, 1.0))).clamp(0.0, 1.0) \ No newline at end of file