working nodes

This commit is contained in:
Waheed-Mousad
2026-07-07 16:43:01 +03:00
commit 074469237f
11 changed files with 879 additions and 0 deletions
+4
View File
@@ -0,0 +1,4 @@
/.idea
/__pycache__
/nodes/__pycache__
/utils/__pycache__
View File
+35
View File
@@ -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"]
View File
+74
View File
@@ -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",
}
+104
View File
@@ -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",
}
+152
View File
@@ -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",
}
+56
View File
@@ -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",
}
+86
View File
@@ -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",
}
View File
+368
View File
@@ -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)