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"}