diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..3e604a6 --- /dev/null +++ b/__init__.py @@ -0,0 +1,45 @@ +""" +@title: Rect +@nickname: Rect +@description: Rectangle selection and utilities for ComfyUI (modular). +""" +import os, sys, pkgutil, importlib, logging +import nodes + +PACK_KEY = "ComfyUI-Rect" # used for front-end assets + +_PACK_DIR = os.path.dirname(os.path.realpath(__file__)) +_PY_DIR = os.path.join(_PACK_DIR, "py") +_JS_DIR = os.path.join(_PACK_DIR, "js") + +# Make ./py importable for dynamic module loading +if _PY_DIR not in sys.path: + sys.path.append(_PY_DIR) + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +def _merge_module(mod): + added = [] + if hasattr(mod, "NODE_CLASS_MAPPINGS"): + NODE_CLASS_MAPPINGS.update(mod.NODE_CLASS_MAPPINGS) + added.extend(mod.NODE_CLASS_MAPPINGS.keys()) + if hasattr(mod, "NODE_DISPLAY_NAME_MAPPINGS"): + NODE_DISPLAY_NAME_MAPPINGS.update(mod.NODE_DISPLAY_NAME_MAPPINGS) + logging.info(f"[Rect] loaded {mod.__name__}: {', '.join(added) or 'no nodes'}") + +# Auto-load all .py files in ./py (except those starting with "_") +for _, modname, ispkg in pkgutil.iter_modules([_PY_DIR]): + if ispkg or modname.startswith("_"): + continue + try: + mod = importlib.import_module(modname) + _merge_module(mod) + except Exception as e: + logging.exception(f"[Rect] failed to load '{modname}': {e}") + +# Serve front-end JS for this pack +if os.path.isdir(_JS_DIR): + nodes.EXTENSION_WEB_DIRS[PACK_KEY] = _JS_DIR + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/js/index.js b/js/index.js new file mode 100644 index 0000000..d35b74c --- /dev/null +++ b/js/index.js @@ -0,0 +1,231 @@ +// ComfyUI-Rect front-end (Rect / Select) — marching ants, image-required, auto-prefill +import { app } from "../../scripts/app.js"; + +function setWidget(node, name, value) { + const w = node.widgets?.find(w => w.name === name); + if (w) { + w.value = value; + node.onWidgetChanged?.(name, value, w); + } + node.properties[name] = value; + node.setDirtyCanvas?.(true, true); +} + +function toast(text, ms = 1800) { + const div = document.createElement("div"); + Object.assign(div.style, { + position: "fixed", right: "16px", bottom: "16px", + background: "rgba(20,20,20,.9)", color: "#eee", + padding: "10px 12px", borderRadius: "10px", + zIndex: 10000, font: "12px/1.3 system-ui, sans-serif", + boxShadow: "0 6px 16px rgba(0,0,0,.35)" + }); + div.textContent = text; + document.body.appendChild(div); + setTimeout(() => div.remove(), ms); +} + +function upstreamFilenameFromImageInput(node) { + const idx = node.inputs?.findIndex(i => i.name === "image"); + if (idx == null || idx < 0) return null; + const linkId = node.inputs[idx]?.link; + if (!linkId) return null; + const link = app.graph.links?.[linkId]; + const upstream = link ? app.graph._nodes_by_id?.[link.origin_id] : null; + const w = upstream?.widgets?.find(w => w.name === "image" && typeof w.value === "string" && w.value.length); + return w ? String(w.value) : null; +} + +function buildInputViewURL(nameOrPath) { + let p = String(nameOrPath).replace(/\\/g, "/"); + const parts = p.split("/"); + const file = parts.pop(); + const subfolder = parts.join("/"); + const ext = (file.split(".").pop() || "png").toLowerCase(); + const format = ext === "jpg" || ext === "jpeg" ? "jpeg" : "png"; + const u = new URL(`${location.origin}/view`); + u.searchParams.set("type", "input"); + u.searchParams.set("filename", file); + u.searchParams.set("subfolder", subfolder); + u.searchParams.set("format", format); + return u.toString(); +} + +function openRectSelect(node) { + const upstreamName = upstreamFilenameFromImageInput(node); + if (!upstreamName) { toast("Rect / Select: connect an image to the 'image' input."); return; } + const upstreamURL = buildInputViewURL(upstreamName); + + // Overlay & modal + const overlay = document.createElement("div"); + Object.assign(overlay.style, { + position: "fixed", inset: "0", background: "rgba(0,0,0,0.6)", + zIndex: 9999, display: "flex", alignItems: "center", justifyContent: "center" + }); + + const modal = document.createElement("div"); + Object.assign(modal.style, { + background: "#111", color: "#eee", padding: "16px", borderRadius: "12px", + width: "min(92vw, 1100px)", maxHeight: "90vh", + display: "grid", gridTemplateRows: "auto 1fr auto", gap: "12px", + boxShadow: "0 10px 30px rgba(0,0,0,0.5)" + }); + overlay.appendChild(modal); + + // Header + const header = document.createElement("div"); + Object.assign(header.style, { display: "flex", justifyContent: "space-between", alignItems: "center" }); + header.innerHTML = `
Rect / Select
`; + const closeBtn = document.createElement("button"); + closeBtn.textContent = "Close"; + closeBtn.onclick = () => document.body.removeChild(overlay); + header.appendChild(closeBtn); + modal.appendChild(header); + + // Canvas area + const area = document.createElement("div"); + Object.assign(area.style, { overflow: "auto", background: "#222", padding: "8px" }); + const canvas = document.createElement("canvas"); + const ctx = canvas.getContext("2d"); + area.appendChild(canvas); + modal.appendChild(area); + + // Footer (coords left, Apply right) + const footer = document.createElement("div"); + Object.assign(footer.style, { display: "flex", gap: "12px", alignItems: "center", flexWrap: "wrap" }); + + const coords = document.createElement("div"); + Object.assign(coords.style, { color: "#aaa", fontStyle: "italic", fontSize: "12px", minHeight: "1em" }); + + const spacer = document.createElement("div"); + spacer.style.flex = "1"; + + const applyBtn = document.createElement("button"); + applyBtn.textContent = "Apply Rect"; + applyBtn.disabled = true; + Object.assign(applyBtn.style, { padding: "10px 14px", fontWeight: "600", borderRadius: "8px", border: "none", cursor: "pointer" }); + + footer.append(coords, spacer, applyBtn); + modal.appendChild(footer); + + // State + let img = new Image(), imgLoaded = false; + let dragging = false, sx = 0, sy = 0, cx = 0, cy = 0; + let antsOffset = 0; + + // Helpers + function clampRect(x, y, w, h, W, H) { + x = Math.max(0, Math.min(x, W)); + y = Math.max(0, Math.min(y, H)); + w = Math.max(1, Math.min(w, W)); + h = Math.max(1, Math.min(h, H)); + if (x + w > W) w = Math.max(1, W - x); + if (y + h > H) h = Math.max(1, H - y); + return [x, y, w, h]; + } + + function setSelectionFromNode() { + if (!imgLoaded) return; + const W = img.naturalWidth, H = img.naturalHeight; + let nx = Number(node.properties?.x ?? 0); + let ny = Number(node.properties?.y ?? 0); + let nw = Number(node.properties?.w ?? Math.floor(W / 2)); + let nh = Number(node.properties?.h ?? Math.floor(H / 2)); + [nx, ny, nw, nh] = clampRect(nx, ny, nw, nh, W, H); + const sxScale = canvas.width / W; + const syScale = canvas.height / H; + sx = Math.round(nx * sxScale); sy = Math.round(ny * syScale); + cx = Math.round((nx + nw) * sxScale); cy = Math.round((ny + nh) * syScale); + draw(); + } + + function fitToViewport() { + if (!imgLoaded) return; + const maxW = Math.min(window.innerWidth * 0.88, 1600); + const maxH = Math.min(window.innerHeight * 0.65, 900); + const scale = Math.min(maxW / img.naturalWidth, maxH / img.naturalHeight, 1); + canvas.width = Math.max(2, Math.floor(img.naturalWidth * scale)); + canvas.height = Math.max(2, Math.floor(img.naturalHeight * scale)); + setSelectionFromNode(); + } + + function draw() { + ctx.fillStyle = "#333"; ctx.fillRect(0, 0, canvas.width, canvas.height); + if (imgLoaded) ctx.drawImage(img, 0, 0, canvas.width, canvas.height); + const x = Math.min(sx, cx), y = Math.min(sy, cy); + const w = Math.abs(cx - sx), h = Math.abs(cy - sy); + if (w > 0 && h > 0) { + const seg = 8; + ctx.lineWidth = 2; + ctx.setLineDash([seg, seg]); + ctx.lineDashOffset = -antsOffset; ctx.strokeStyle = "#fff"; ctx.strokeRect(x, y, w, h); + ctx.lineDashOffset = seg - antsOffset; ctx.strokeStyle = "#000"; ctx.strokeRect(x, y, w, h); + coords.textContent = `x=${x}, y=${y}, w=${w}, h=${h}`; + } else { + coords.textContent = imgLoaded ? "Drag to draw a rectangle." : ""; + } + } + + function animate() { + antsOffset = (antsOffset + 1) % 16; + draw(); + if (document.body.contains(overlay)) requestAnimationFrame(animate); + } + + function loadFromURL(url) { + img = new Image(); + img.crossOrigin = "anonymous"; + img.onload = () => { imgLoaded = true; applyBtn.disabled = false; fitToViewport(); }; + img.onerror = () => { toast("Rect / Select: could not load upstream image."); }; + img.src = url; + } + + // Interactions + canvas.addEventListener("mousedown", (e) => { + if (!imgLoaded) return; + const r = canvas.getBoundingClientRect(); + sx = cx = e.clientX - r.left; sy = cy = e.clientY - r.top; dragging = true; draw(); + }); + canvas.addEventListener("mousemove", (e) => { + if (!dragging || !imgLoaded) return; + const r = canvas.getBoundingClientRect(); + cx = e.clientX - r.left; cy = e.clientY - r.top; draw(); + }); + window.addEventListener("mouseup", () => { if (dragging) { dragging = false; draw(); } }); + window.addEventListener("resize", fitToViewport); + + applyBtn.onclick = () => { + if (!imgLoaded) return; + const scaleX = img.naturalWidth / canvas.width; + const scaleY = img.naturalHeight / canvas.height; + const x = Math.round(Math.min(sx, cx) * scaleX); + const y = Math.round(Math.min(sy, cy) * scaleY); + const w = Math.round(Math.abs(cx - sx) * scaleX); + const h = Math.round(Math.abs(cy - sy) * scaleY); + if (w < 1 || h < 1) { toast("Draw a rectangle first."); return; } + setWidget(node, "x", x); setWidget(node, "y", y); + setWidget(node, "w", w); setWidget(node, "h", h); + setTimeout(() => document.body.contains(overlay) && document.body.removeChild(overlay), 250); + }; + + document.body.appendChild(overlay); + requestAnimationFrame(animate); + loadFromURL(upstreamURL); +} + +// Register: attach button to RectSelect nodes +app.registerExtension({ + name: "ComfyUI-Rect", + nodeCreated(node) { + if (node?.comfyClass !== "RectSelect") return; + if (node.widgets?.some(w => w.__rect_btn)) return; + const btn = node.addWidget("button", "Open Rect / Select", "open", () => openRectSelect(node)); + btn.__rect_btn = true; + + node.properties ??= {}; + node.properties.x ??= 0; node.properties.y ??= 0; + node.properties.w ??= 256; node.properties.h ??= 256; + + console.log("[Rect] button attached to node", node.id); + }, +}); diff --git a/py/__init__.py b/py/__init__.py new file mode 100644 index 0000000..0f8c0e6 --- /dev/null +++ b/py/__init__.py @@ -0,0 +1 @@ +# (intentionally empty) diff --git a/py/__pycache__/rect_crop.cpython-311.pyc b/py/__pycache__/rect_crop.cpython-311.pyc new file mode 100644 index 0000000..2fc34fa Binary files /dev/null and b/py/__pycache__/rect_crop.cpython-311.pyc differ diff --git a/py/__pycache__/rect_fill.cpython-311.pyc b/py/__pycache__/rect_fill.cpython-311.pyc new file mode 100644 index 0000000..035fb3e Binary files /dev/null and b/py/__pycache__/rect_fill.cpython-311.pyc differ diff --git a/py/__pycache__/rect_mask.cpython-311.pyc b/py/__pycache__/rect_mask.cpython-311.pyc new file mode 100644 index 0000000..dd300f7 Binary files /dev/null and b/py/__pycache__/rect_mask.cpython-311.pyc differ diff --git a/py/__pycache__/rect_select.cpython-311.pyc b/py/__pycache__/rect_select.cpython-311.pyc new file mode 100644 index 0000000..2b5f3bb Binary files /dev/null and b/py/__pycache__/rect_select.cpython-311.pyc differ diff --git a/py/rect_crop.py b/py/rect_crop.py new file mode 100644 index 0000000..9b0ffb4 --- /dev/null +++ b/py/rect_crop.py @@ -0,0 +1,65 @@ +# RectCrop node (display: "Rect / Crop") +# Crops an IMAGE to the given RECT (x,y,w,h in pixels). +import torch + +def _image_size(image): + if isinstance(image, torch.Tensor): + if image.dim() == 4: # [B,H,W,C] + return int(image.shape[2]), int(image.shape[1]) + if image.dim() == 3: # [H,W,C] + return int(image.shape[1]), int(image.shape[0]) + return 512, 512 + +def _clamp_rect_for_crop(x, y, w, h, W, H): + # Clamp top-left *inside* the image so slicing never returns empty. + if W <= 0 or H <= 0: + return 0, 0, 1, 1 + x = max(0, min(int(x), W - 1)) + y = max(0, min(int(y), H - 1)) + # Width/height must fit within the remaining bounds from (x,y) + w = max(1, min(int(w), W - x)) + h = max(1, min(int(h), H - y)) + return x, y, w, h + +class RectCrop: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "rect": ("RECT",), # {"x":int,"y":int,"w":int,"h":int} + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + FUNCTION = "run" + CATEGORY = "Rect" + + def run(self, image, rect): + # Extract rect safely + try: + x = int(rect.get("x", 0)) + y = int(rect.get("y", 0)) + w = int(rect.get("w", 1)) + h = int(rect.get("h", 1)) + except Exception: + x, y, w, h = 0, 0, 1, 1 + + W, H = _image_size(image) + x, y, w, h = _clamp_rect_for_crop(x, y, w, h, W, H) + + if not isinstance(image, torch.Tensor): + raise ValueError("RectCrop: expected torch.Tensor IMAGE") + + if image.dim() == 4: # [B,H,W,C] + cropped = image[:, y:y+h, x:x+w, :] + elif image.dim() == 3: # [H,W,C] + cropped = image[y:y+h, x:x+w, :] + else: + raise ValueError(f"RectCrop: unsupported IMAGE dims {image.shape}") + + return (cropped,) + +NODE_CLASS_MAPPINGS = {"RectCrop": RectCrop} +NODE_DISPLAY_NAME_MAPPINGS = {"RectCrop": "Rect / Crop"} diff --git a/py/rect_fill.py b/py/rect_fill.py new file mode 100644 index 0000000..b51f3f5 --- /dev/null +++ b/py/rect_fill.py @@ -0,0 +1,126 @@ +# Rect / Fill — fill or outline a RECT region on IMAGE with color & opacity (optional feather) +import torch +import torch.nn.functional as F + +def _image_size(image): + if isinstance(image, torch.Tensor): + if image.dim() == 4: # [B,H,W,C] + return int(image.shape[2]), int(image.shape[1]) + if image.dim() == 3: # [H,W,C] + return int(image.shape[1]), int(image.shape[0]) + return 512, 512 + +def _clamp_rect(x, y, w, h, W, H): + x = max(0, min(int(x), max(0, W - 1))) + y = max(0, min(int(y), max(0, H - 1))) + w = max(1, min(int(w), W - x)) + h = max(1, min(int(h), H - y)) + return x, y, w, h + +def _gaussian_kernel1d(radius, sigma, device): + xs = torch.arange(-radius, radius + 1, device=device, dtype=torch.float32) + k = torch.exp(-(xs**2) / (2 * sigma * sigma)) + k /= k.sum().clamp_min(1e-8) + return k + +def _gaussian_blur(mask, radius): + if radius < 1: + return mask + B, H, W = mask.shape + device = mask.device + sigma = max(0.5, radius / 2.5) + k1d = _gaussian_kernel1d(radius, sigma, device) + x = mask.unsqueeze(1) # [B,1,H,W] + kh = k1d.view(1, 1, 1, -1) + kv = k1d.view(1, 1, -1, 1) + x = F.pad(x, (radius, radius, 0, 0), mode="reflect") + x = F.conv2d(x, kh) + x = F.pad(x, (0, 0, radius, radius), mode="reflect") + x = F.conv2d(x, kv) + return x.squeeze(1).clamp(0.0, 1.0) + +class RectFill: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "rect": ("RECT",), + "r": ("INT", {"default": 255, "min": 0, "max": 255}), + "g": ("INT", {"default": 0, "min": 0, "max": 255}), + "b": ("INT", {"default": 0, "min": 0, "max": 255}), + "opacity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0}), + "mode": ("STRING", {"default": "fill", "choices": ["fill", "outline"]}), + "thickness": ("INT", {"default": 4, "min": 1, "max": 1024}), + "feather": ("INT", {"default": 0, "min": 0, "max": 256}), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + FUNCTION = "run" + CATEGORY = "Rect" + + def run(self, image, rect, r, g, b, opacity, mode, thickness, feather): + if not isinstance(image, torch.Tensor): + raise ValueError("RectFill: expected torch.Tensor IMAGE") + + # Parse rect + try: + x = int(rect.get("x", 0)); y = int(rect.get("y", 0)) + w = int(rect.get("w", 1)); h = int(rect.get("h", 1)) + except Exception: + x, y, w, h = 0, 0, 1, 1 + + # Shapes + if image.dim() == 4: + B, H, W, C = int(image.shape[0]), int(image.shape[1]), int(image.shape[2]), int(image.shape[3]) + img = image + elif image.dim() == 3: + B, H, W, C = 1, int(image.shape[0]), int(image.shape[1]), int(image.shape[2]) + img = image.unsqueeze(0) # [1,H,W,C] + else: + raise ValueError(f"RectFill: unsupported IMAGE dims {tuple(image.shape)}") + + device = img.device + x, y, w, h = _clamp_rect(x, y, w, h, W, H) + + # Build alpha mask in [B,H,W] + alpha = torch.zeros((B, H, W), device=device, dtype=torch.float32) + + if mode == "fill": + alpha[:, y:y+h, x:x+w] = 1.0 + else: # outline + # Outer rect + alpha[:, y:y+h, x:x+w] = 1.0 + # Inner rect to subtract + inner_w = max(0, w - 2 * thickness) + inner_h = max(0, h - 2 * thickness) + if inner_w > 0 and inner_h > 0: + ix = x + thickness + iy = y + thickness + alpha[:, iy:iy+inner_h, ix:ix+inner_w] = 0.0 + + # Feather (Gaussian blur) + if feather > 0: + alpha = _gaussian_blur(alpha, int(feather)) + + # Apply opacity + alpha = (alpha * float(opacity)).clamp(0.0, 1.0) + + # Color tensor [B,1,1,3] in 0..1 + color = torch.tensor([r, g, b], device=device, dtype=torch.float32) / 255.0 + color = color.view(1, 1, 1, 3).expand(B, 1, 1, 3) + + # Blend: out = alpha*color + (1-alpha)*img + alpha4 = alpha.unsqueeze(-1) # [B,H,W,1] + out = (alpha4 * color) + ((1.0 - alpha4) * img) + out = out.clamp(0.0, 1.0) + + if image.dim() == 3: + out = out.squeeze(0) + + return (out,) + +NODE_CLASS_MAPPINGS = {"RectFill": RectFill} +NODE_DISPLAY_NAME_MAPPINGS = {"RectFill": "Rect / Fill"} diff --git a/py/rect_mask.py b/py/rect_mask.py new file mode 100644 index 0000000..29d073c --- /dev/null +++ b/py/rect_mask.py @@ -0,0 +1,141 @@ +# Rect / Mask — build a MASK from a RECT, with optional feather/invert/combine +import math +import torch +import torch.nn.functional as F + +def _image_size(image): + if isinstance(image, torch.Tensor): + if image.dim() == 4: # [B,H,W,C] + return int(image.shape[2]), int(image.shape[1]) + if image.dim() == 3: # [H,W,C] + return int(image.shape[1]), int(image.shape[0]) + return 512, 512 + +def _clamp_rect(x, y, w, h, W, H): + x = max(0, min(int(x), max(0, W - 1))) + y = max(0, min(int(y), max(0, H - 1))) + w = max(1, min(int(w), W - x)) + h = max(1, min(int(h), H - y)) + return x, y, w, h + +def _ensure_mask_shape(mask, B, H, W, device): + # Accept [H,W], [B,H,W], or [B,1,H,1] quirky shapes from some packs + if mask is None: + return None + if mask.dim() == 2: + mask = mask.unsqueeze(0) # [1,H,W] + if mask.dim() == 4 and mask.shape[1] == 1 and mask.shape[3] == 1: + mask = mask[:, 0, :, :] # [B,H,W] + if mask.dim() != 3: + raise ValueError(f"RectMask: unsupported mask shape {tuple(mask.shape)}") + # Broadcast or clamp batch as needed + if mask.shape[0] == 1 and B > 1: + mask = mask.expand(B, H, W).clone() + elif mask.shape[0] != B: + # If sizes mismatch, just take first and broadcast + mask = mask[:1].expand(B, H, W).clone() + # Resize if spatial size mismatches + if (mask.shape[1] != H) or (mask.shape[2] != W): + mask = F.interpolate(mask.unsqueeze(1), size=(H, W), mode="bilinear", align_corners=False).squeeze(1) + return mask.to(device=device, dtype=torch.float32).clamp(0.0, 1.0) + +def _gaussian_kernel1d(radius, sigma, device): + # radius: pixels; kernel size = 2*radius+1 + xs = torch.arange(-radius, radius + 1, device=device, dtype=torch.float32) + k = torch.exp(-(xs**2) / (2 * sigma * sigma)) + k /= k.sum().clamp_min(1e-8) + return k + +def _gaussian_blur(mask, radius): + # mask: [B,H,W], radius >= 1 + if radius < 1: + return mask + B, H, W = mask.shape + device = mask.device + sigma = max(0.5, radius / 2.5) + k1d = _gaussian_kernel1d(radius, sigma, device) + # separable blur: first horizontal, then vertical + x = mask.unsqueeze(1) # [B,1,H,W] + kh = k1d.view(1, 1, 1, -1) + kv = k1d.view(1, 1, -1, 1) + x = F.pad(x, (radius, radius, 0, 0), mode="reflect") + x = F.conv2d(x, kh) + x = F.pad(x, (0, 0, radius, radius), mode="reflect") + x = F.conv2d(x, kv) + return x.squeeze(1).clamp(0.0, 1.0) + +class RectMask: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "rect": ("RECT",), # {"x","y","w","h"} + "feather": ("INT", {"default": 0, "min": 0, "max": 256}), + "invert": ("BOOLEAN", {"default": False}), + "combine": ("STRING", {"default": "replace", + "choices": ["replace", "union", "intersect", "subtract", "multiply"]}), + }, + "optional": { + "existing_mask": ("MASK",), + } + } + + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("mask",) + FUNCTION = "run" + CATEGORY = "Rect" + + def run(self, image, rect, feather, invert, combine, existing_mask=None): + if not isinstance(image, torch.Tensor): + raise ValueError("RectMask: expected torch.Tensor IMAGE") + + # Parse rect + try: + x = int(rect.get("x", 0)); y = int(rect.get("y", 0)) + w = int(rect.get("w", 1)); h = int(rect.get("h", 1)) + except Exception: + x, y, w, h = 0, 0, 1, 1 + + # Get sizes and clamp rect + if image.dim() == 4: + B, H, W, C = int(image.shape[0]), int(image.shape[1]), int(image.shape[2]), int(image.shape[3]) + elif image.dim() == 3: + B, H, W, C = 1, int(image.shape[0]), int(image.shape[1]), int(image.shape[2]) + else: + raise ValueError(f"RectMask: unsupported IMAGE dims {tuple(image.shape)}") + + device = image.device + x, y, w, h = _clamp_rect(x, y, w, h, W, H) + + # Build binary rect mask + mask = torch.zeros((B, H, W), device=device, dtype=torch.float32) + mask[:, y:y+h, x:x+w] = 1.0 + + # Feather (Gaussian) + if feather > 0: + radius = int(feather) + mask = _gaussian_blur(mask, radius) + + # Invert + if invert: + mask = 1.0 - mask + + # Combine with existing_mask + if existing_mask is not None: + em = _ensure_mask_shape(existing_mask, B, H, W, device) + if combine == "replace": + mask = mask + elif combine == "union": + mask = torch.maximum(em, mask) + elif combine == "intersect": + mask = torch.minimum(em, mask) + elif combine == "subtract": + mask = (em - mask).clamp(0.0, 1.0) + elif combine == "multiply": + mask = (em * mask).clamp(0.0, 1.0) + + return (mask,) + +NODE_CLASS_MAPPINGS = {"RectMask": RectMask} +NODE_DISPLAY_NAME_MAPPINGS = {"RectMask": "Rect / Mask"} diff --git a/py/rect_select.py b/py/rect_select.py new file mode 100644 index 0000000..93bb336 --- /dev/null +++ b/py/rect_select.py @@ -0,0 +1,47 @@ +# RectSelect node (display: "Rect / Select") +import torch + +def _image_size(image): + if isinstance(image, torch.Tensor): + if image.dim() == 4: # [B,H,W,C] + return int(image.shape[2]), int(image.shape[1]) + if image.dim() == 3: # [H,W,C] + return int(image.shape[1]), int(image.shape[0]) + return 512, 512 + +def _clamp_rect(x, y, w, h, W, H): + x = max(0, min(int(x), max(0, W))) + y = max(0, min(int(y), max(0, H))) + w = max(1, int(w)) + h = max(1, int(h)) + if x + w > W: w = max(1, W - x) + if y + h > H: h = max(1, H - y) + return x, y, w, h + +class RectSelect: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "x": ("INT", {"default": 0, "min": 0}), + "y": ("INT", {"default": 0, "min": 0}), + "w": ("INT", {"default": 256, "min": 1}), + "h": ("INT", {"default": 256, "min": 1}), + } + } + + # Output a RECT object and the four ints (compat with existing crop nodes) + RETURN_TYPES = ("RECT", "INT", "INT", "INT", "INT") + RETURN_NAMES = ("rect", "x", "y", "w", "h") + FUNCTION = "run" + CATEGORY = "Rect" + + def run(self, image, x, y, w, h): + W, H = _image_size(image) + x, y, w, h = _clamp_rect(x, y, w, h, W, H) + rect = {"x": x, "y": y, "w": w, "h": h} + return (rect, x, y, w, h) + +NODE_CLASS_MAPPINGS = {"RectSelect": RectSelect} +NODE_DISPLAY_NAME_MAPPINGS = {"RectSelect": "Rect / Select"}