Files
INuBq8-ComfyUI-MultiMaskOps/utils/mask_ops.py
T
2026-07-07 16:43:01 +03:00

368 lines
13 KiB
Python

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)