From ff86e4009aadf13025c0affaa8991c9a26fd89fd Mon Sep 17 00:00:00 2001 From: Marco Date: Tue, 14 Apr 2026 00:39:27 -0300 Subject: [PATCH] Add SAM2 adapter nodes and interactive point collector --- __init__.py | 14 +++ image_processing/point_collector.py | 60 ++++++++++++ image_processing/segs_adapter.py | 147 ++++++++++++++++++++++++++++ js/point_collector.js | 134 +++++++++++++++++++++++++ 4 files changed, 355 insertions(+) create mode 100644 image_processing/point_collector.py create mode 100644 image_processing/segs_adapter.py create mode 100644 js/point_collector.js diff --git a/__init__.py b/__init__.py index 77399c5..7c49639 100644 --- a/__init__.py +++ b/__init__.py @@ -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 diff --git a/image_processing/point_collector.py b/image_processing/point_collector.py new file mode 100644 index 0000000..621862a --- /dev/null +++ b/image_processing/point_collector.py @@ -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 diff --git a/image_processing/segs_adapter.py b/image_processing/segs_adapter.py new file mode 100644 index 0000000..61eec06 --- /dev/null +++ b/image_processing/segs_adapter.py @@ -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),) diff --git a/js/point_collector.js b/js/point_collector.js new file mode 100644 index 0000000..a2a8295 --- /dev/null +++ b/js/point_collector.js @@ -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; + }; + } + } +});