Add files via upload
This commit is contained in:
+45
@@ -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"]
|
||||
+231
@@ -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 = `<div style="font-weight:600">Rect / Select</div>`;
|
||||
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);
|
||||
},
|
||||
});
|
||||
@@ -0,0 +1 @@
|
||||
# (intentionally empty)
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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"}
|
||||
+126
@@ -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"}
|
||||
+141
@@ -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"}
|
||||
@@ -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"}
|
||||
Reference in New Issue
Block a user