Add SAM2 adapter nodes and interactive point collector

This commit is contained in:
Marco
2026-04-14 00:39:27 -03:00
parent c85709b8af
commit ff86e4009a
4 changed files with 355 additions and 0 deletions
+14
View File
@@ -51,6 +51,8 @@ from .video.ltxv_vid2vid import LTXVVid2Vid
from .loaders.folder_image_loader import FolderImageLoader
from .logic_management.dataset_loader import DatasetLoader
from .logic_management.image_list_sampler import ImageListSampler
from .image_processing.segs_adapter import SEGStoBBox, SEGStoSAM2Points, GetFirstFrame, ManualPointToSAM2, RefineMask
from .image_processing.point_collector import PointCollectorSAM2
from .core import server_routes # Register Custom API Routes
NODE_CLASS_MAPPINGS = {
@@ -93,6 +95,12 @@ NODE_CLASS_MAPPINGS = {
"ImageListSampler": ImageListSampler,
"LTXVMultiGuide": LTXVMultiGuide,
"LTXVVid2Vid": LTXVVid2Vid,
"SEGStoBBox": SEGStoBBox,
"SEGStoSAM2Points": SEGStoSAM2Points,
"GetFirstFrame": GetFirstFrame,
"ManualPointToSAM2": ManualPointToSAM2,
"RefineMask": RefineMask,
"PointCollectorSAM2": PointCollectorSAM2,
}
# Add V3 nodes if the new API is available
@@ -138,6 +146,12 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ImageListSampler": "Image List Sampler",
"LTXVMultiGuide": "LTXV Multi Guide (N Frames)",
"LTXVVid2Vid": "LTXV Vid2Vid Encode",
"SEGStoBBox": "SEGS to BBox",
"SEGStoSAM2Points": "SEGS to SAM2 Points (JSON)",
"GetFirstFrame": "Get First Frame (Batch to Single)",
"ManualPointToSAM2": "Manual Point to SAM2 (JSON)",
"RefineMask": "Refine Mask (Expand & Blur)",
"PointCollectorSAM2": "Interactive Point Collector (SAM2)",
}
# Add V3 display names if available
+60
View File
@@ -0,0 +1,60 @@
import torch
import numpy as np
import json
import io
import base64
from PIL import Image
class PointCollectorSAM2:
"""
Interactive Point Collector for SAM 2.
Outputs pixel coordinates in JSON format: [{"x": 1, "y": 2}]
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
# Hidden widgets populated by JS
"coordinates": ("STRING", {"multiline": False, "default": "[]"}),
"neg_coordinates": ("STRING", {"multiline": False, "default": "[]"}),
},
}
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("pos_points_json", "neg_points_json")
FUNCTION = "collect"
CATEGORY = "AnotherUtils/sam2"
OUTPUT_NODE = True
def collect(self, image, coordinates, neg_coordinates):
# Coordinates from JS are already in pixel units relative to the image
# We just need to ensure they are valid JSON strings for our SAM2 nodes
pos_json = coordinates if coordinates and coordinates.strip() else "[]"
neg_json = neg_coordinates if neg_coordinates and neg_coordinates.strip() else "[]"
# Send image to the JS widget via a UI message
img_base64 = self.tensor_to_base64(image)
return {
"ui": {"bg_image": [img_base64]},
"result": (pos_json, neg_json)
}
def tensor_to_base64(self, tensor):
# Convert from [B, H, W, C] to PIL Image (first frame)
img_array = tensor[0].cpu().numpy()
img_array = (img_array * 255).astype(np.uint8)
pil_img = Image.fromarray(img_array)
buffered = io.BytesIO()
pil_img.save(buffered, format="JPEG", quality=75)
img_bytes = buffered.getvalue()
img_base64 = base64.b64encode(img_bytes).decode('utf-8')
return img_base64
@classmethod
def IS_CHANGED(cls, image, coordinates, neg_coordinates):
return float("nan") # Always update UI when points change
+147
View File
@@ -0,0 +1,147 @@
import numpy as np
import json
class SEGStoBBox:
@classmethod
def INPUT_TYPES(s):
return {"required": {"segs": ("SEGS",),}}
RETURN_TYPES = ("BBOX",)
FUNCTION = "convert"
CATEGORY = "AnotherUtils/SEGS"
def convert(self, segs):
"""
Converts Impact Pack SEGS to SAM 2 BBOX format.
SAM 2 expects a list of lists of bounding boxes: [ [ [x1,y1,x2,y2], ... ] ]
"""
if not segs or len(segs) < 2:
return ([[]],)
bboxes = []
for seg in segs[1]:
# seg.bbox is (x1, y1, x2, y2)
bboxes.append(list(seg.bbox))
# Return as a batch of boxes (standard for single image context)
return ([bboxes],)
class SEGStoSAM2Points:
@classmethod
def INPUT_TYPES(s):
return {"required": {"segs": ("SEGS",),}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("json_points",)
FUNCTION = "convert"
CATEGORY = "AnotherUtils/SEGS"
def convert(self, segs):
"""
Converts Impact Pack SEGS to SAM 2 JSON point format.
Useful for 'coordinates_positive' input in SAM 2 Video nodes.
"""
if not segs or len(segs) < 2 or not segs[1]:
print("!!! [AnotherUtils] YOLO found no objects. SAM2 tracking might fail.")
return ("[]",)
points = []
for seg in segs[1]:
x1, y1, x2, y2 = seg.bbox
center_x = int((x1 + x2) / 2)
center_y = int((y1 + y2) / 2)
points.append({"x": center_x, "y": center_y})
return (json.dumps(points),)
class GetFirstFrame:
@classmethod
def INPUT_TYPES(s):
return {"required": {"images": ("IMAGE",),}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "get"
CATEGORY = "AnotherUtils/image"
def get(self, images):
return (images[0:1],)
class ManualPointToSAM2:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"x": ("INT", {"default": 333, "min": 0, "max": 8192}),
"y": ("INT", {"default": 333, "min": 0, "max": 8192}),
"num_points": ("INT", {"default": 1, "min": 1, "max": 100}),
"radius": ("INT", {"default": 0, "min": 0, "max": 500}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "convert"
CATEGORY = "AnotherUtils/image"
def convert(self, x, y, num_points, radius, seed):
import json
import random
import math
random.seed(seed)
points = []
if num_points <= 1 or radius <= 0:
points.append({"x": x, "y": y})
else:
for _ in range(num_points):
# Use sqrt(r) to get uniform distribution in a circle
r = radius * math.sqrt(random.random())
theta = random.random() * 2 * math.pi
px = int(x + r * math.cos(theta))
py = int(y + r * math.sin(theta))
points.append({"x": px, "y": py})
return (json.dumps(points),)
class RefineMask:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask": ("MASK",),
"expand": ("INT", {"default": 0, "min": -100, "max": 100, "step": 1}),
"blur": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}),
}
}
RETURN_TYPES = ("MASK",)
FUNCTION = "refine"
CATEGORY = "AnotherUtils/mask"
def refine(self, mask, expand, blur):
import scipy.ndimage
import numpy as np
import torch
# mask shape is [B, H, W] or [H, W]
mask_np = mask.cpu().numpy()
# Dilate/Erode
if expand != 0:
if expand > 0:
mask_np = scipy.ndimage.binary_dilation(mask_np, iterations=expand)
else:
mask_np = scipy.ndimage.binary_erosion(mask_np, iterations=abs(expand))
mask_np = mask_np.astype(np.float32)
# Blur
if blur > 0:
for i in range(mask_np.shape[0] if len(mask_np.shape) == 3 else 1):
m = mask_np[i] if len(mask_np.shape) == 3 else mask_np
m = scipy.ndimage.gaussian_filter(m, sigma=blur)
if len(mask_np.shape) == 3: mask_np[i] = m
else: mask_np = m
return (torch.from_numpy(mask_np),)
+134
View File
@@ -0,0 +1,134 @@
import { app } from "../../scripts/app.js";
// Helper to hide sync widgets
function hideWidget(node, widget) {
if (!widget) return;
widget.computeSize = () => [0, -4];
widget.type = "converted-widget";
if (widget.element) {
widget.element.style.display = "none";
}
}
console.log("[PointCollector] Loading PointCollector extension...");
app.registerExtension({
name: "AnotherUtils.PointCollector",
async beforeRegisterNodeDef(nodeType, nodeData) {
if (nodeData.name === "PointCollectorSAM2") {
const onNodeCreated = nodeType.prototype.onNodeCreated;
nodeType.prototype.onNodeCreated = function () {
console.log("[PointCollector] Node created:", this.id);
const result = onNodeCreated?.apply(this, arguments);
// Initial size
this.size = [400, 400];
// Create container
const container = document.createElement("div");
container.style.cssText = "position: relative; width: 100%; background: #111; display: flex; align-items: center; justify-content: center;";
// Info bar
const infoBar = document.createElement("div");
infoBar.style.cssText = "position: absolute; top: 5px; left: 5px; right: 5px; z-index: 10; display: flex; justify-content: space-between;";
container.appendChild(infoBar);
const counter = document.createElement("div");
counter.style.cssText = "background: rgba(0,0,0,0.8); color: #0f0; padding: 2px 8px; font-size: 11px; font-family: monospace; border-radius: 3px;";
counter.textContent = "Points: 0p / 0n";
infoBar.appendChild(counter);
const btn = document.createElement("button");
btn.textContent = "Clear";
btn.style.cssText = "background: #622; color: #fff; border: none; padding: 2px 8px; cursor: pointer; border-radius: 3px;";
btn.onclick = (e) => {
e.preventDefault();
this.canvasWidget.pos = [];
this.canvasWidget.neg = [];
this.updatePoints();
this.draw();
};
infoBar.appendChild(btn);
// Canvas
const canvas = document.createElement("canvas");
canvas.style.cssText = "display: block; max-width: 100%; max-height: 100%; cursor: crosshair;";
container.appendChild(canvas);
this.canvasWidget = {
canvas, ctx: canvas.getContext("2d"),
pos: [], neg: [], image: null, counter
};
const widget = this.addDOMWidget("canvas", "pointsPreview", container);
widget.computeSize = (width) => [width, this.canvasWidget.height || 300];
console.log("[PointCollector] DOM widget added");
// Hide strings
const wCoords = this.widgets?.find(w => w.name === "coordinates");
const wNegCoords = this.widgets?.find(w => w.name === "neg_coordinates");
console.log("[PointCollector] Widgets found to hide:", { wCoords, wNegCoords });
hideWidget(this, wCoords);
hideWidget(this, wNegCoords);
// Click Handlers
canvas.addEventListener("mousedown", (e) => {
const rect = canvas.getBoundingClientRect();
const scaleX = canvas.width / rect.width;
const scaleY = canvas.height / rect.height;
const x = (e.clientX - rect.left) * scaleX;
const y = (e.clientY - rect.top) * scaleY;
if (e.button === 0 && !e.shiftKey) {
this.canvasWidget.pos.push({x, y});
} else {
this.canvasWidget.neg.push({x, y});
}
this.updatePoints();
this.draw();
});
canvas.addEventListener("contextmenu", (e) => e.preventDefault());
this.onExecuted = (msg) => {
if (msg.bg_image?.[0]) {
const img = new Image();
img.onload = () => {
this.canvasWidget.image = img;
canvas.width = img.width;
canvas.height = img.height;
const aspectRatio = img.height / img.width;
this.canvasWidget.height = (this.size[0] - 20) * aspectRatio;
this.draw();
};
img.src = "data:image/jpeg;base64," + msg.bg_image[0];
}
};
this.updatePoints = () => {
const wPos = this.widgets.find(w => w.name === "coordinates");
const wNeg = this.widgets.find(w => w.name === "neg_coordinates");
if (wPos) wPos.value = JSON.stringify(this.canvasWidget.pos);
if (wNeg) wNeg.value = JSON.stringify(this.canvasWidget.neg);
this.canvasWidget.counter.textContent = `Points: ${this.canvasWidget.pos.length}p / ${this.canvasWidget.neg.length}n`;
};
this.draw = () => {
const {canvas, ctx, image, pos, neg} = this.canvasWidget;
ctx.clearRect(0,0, canvas.width, canvas.height);
if (image) ctx.drawImage(image, 0, 0);
ctx.lineWidth = 2;
pos.forEach(p => {
ctx.fillStyle = "#0f0"; ctx.beginPath(); ctx.arc(p.x, p.y, 6, 0, Math.PI*2); ctx.fill();
});
neg.forEach(p => {
ctx.fillStyle = "#f00"; ctx.beginPath(); ctx.arc(p.x, p.y, 6, 0, Math.PI*2); ctx.fill();
});
};
return result;
};
}
}
});