Add SAM2 adapter nodes and interactive point collector
This commit is contained in:
+14
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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),)
|
||||
@@ -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;
|
||||
};
|
||||
}
|
||||
}
|
||||
});
|
||||
Reference in New Issue
Block a user