368 lines
13 KiB
Python
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) |