working nodes
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
/.idea
|
||||
/__pycache__
|
||||
/nodes/__pycache__
|
||||
/utils/__pycache__
|
||||
+35
@@ -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"]
|
||||
@@ -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
@@ -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
@@ -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",
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user