[Nodes][Added] ResizeMask also synced with Kijai 1.1.7

This commit is contained in:
Salvador E. Tropea
2025-10-06 13:36:24 -03:00
parent d99c165f3d
commit f05c36f520
+304 -88
View File
@@ -2,9 +2,12 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnologïa Industrial
# License: GPL-3.0
# ImagePad, ImageResize are from Kijai (https://github.com/kijai/ComfyUI-KJNodes/)
#
# Project: ComfyUI-ImageMisc
# From code generated by Gemini 2.5 Pro
# Credits:
# - ImagePad, ImageResize and ResizeMask are from Kijai (https://github.com/kijai/ComfyUI-KJNodes/) v1.1.7
# - Assisted by Gemini 2.5 Pro
from copy import deepcopy
import numpy as np
import os
from PIL import Image # Import the Python Imaging Library
@@ -18,7 +21,6 @@ from seconohe.color import color_to_rgb_uint8
from . import main_logger
import torch
import torch.nn.functional as F
from torchvision import transforms
import torchvision.transforms.functional as TF
from typing import Optional
try:
@@ -63,17 +65,20 @@ COLOR_OPT = ("STRING", {
"tooltip": "Color for fill.\n"
"Can be an hexadecimal (#RRGGBB).\n"
"Can comma separated RGB values in [0-255] or [0-1.0] range."})
DEFAULT_UPSCALE = 'bicubic' # transforms.InterpolationMode.BICUBIC.value
UPSCALE_OPT = (ImageScale.upscale_methods, { # [mode.value for mode in transforms.InterpolationMode]
DEFAULT_UPSCALE = 'bicubic' # transforms.InterpolationMode.BICUBIC.value
MASK_UPSCALE = 'nearest-exact' # transforms.InterpolationMode.NEAREST_EXACT.value
BEST_UPSCALE = 'lanczos' # transforms.InterpolationMode.LANCZOS.value
UPSCALE_OPT = (ImageScale.upscale_methods, { # [mode.value for mode in transforms.InterpolationMode]
"default": DEFAULT_UPSCALE,
"tooltip": "Interpolation method for image resize"
})
UPSCALE_OPT_MASK = deepcopy(UPSCALE_OPT)
UPSCALE_OPT_MASK[1]["default"] = MASK_UPSCALE
PAD_SIZE_OPT = ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1, })
SIZE_OPT = ("INT", {"default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 1, })
SIZE_OPT_FI = tuple(SIZE_OPT)
SIZE_OPT = ("INT", {"default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 1})
SIZE_OPT_FI = deepcopy(SIZE_OPT)
SIZE_OPT_FI[1]["forceInput"] = True
MASK_UPSCALE = 'nearest-exact' # transforms.InterpolationMode.NEAREST_EXACT.value
BEST_UPSCALE = 'lanczos' # transforms.InterpolationMode.LANCZOS.value
SIZE_OPT[1]["tooltip"] = "Used when no `get_image_size` is provided"
def tensor_to_pil(tensor: torch.Tensor) -> Image.Image:
@@ -712,7 +717,7 @@ class ImagePad:
"top": PAD_SIZE_OPT,
"bottom": PAD_SIZE_OPT,
"extra_padding": PAD_SIZE_OPT,
"pad_mode": (["edge", "color"],),
"pad_mode": (["edge", "edge_pixel", "color", "pillarbox_blur"],),
"color": COLOR_OPT,
},
"optional": {
@@ -763,27 +768,102 @@ class ImagePad:
padded_width = W + pad_left + pad_right
padded_height = H + pad_top + pad_bottom
out_image = torch.zeros((B, padded_height, padded_width, C), dtype=image.dtype, device=image.device)
# Fill padded areas
# Pillarbox blur mode
if pad_mode == "pillarbox_blur":
def _gaussian_blur_nchw(img_nchw, sigma_px):
if sigma_px <= 0:
return img_nchw
radius = max(1, int(3.0 * float(sigma_px)))
k = 2 * radius + 1
x = torch.arange(-radius, radius + 1, device=img_nchw.device, dtype=img_nchw.dtype)
k1 = torch.exp(-(x * x) / (2.0 * float(sigma_px) * float(sigma_px)))
k1 = k1 / k1.sum()
kx = k1.view(1, 1, 1, k)
ky = k1.view(1, 1, k, 1)
c = img_nchw.shape[1]
kx = kx.repeat(c, 1, 1, 1)
ky = ky.repeat(c, 1, 1, 1)
img_nchw = F.conv2d(img_nchw, kx, padding=(0, radius), groups=c)
img_nchw = F.conv2d(img_nchw, ky, padding=(radius, 0), groups=c)
return img_nchw
out_image = torch.zeros((B, padded_height, padded_width, C), dtype=image.dtype, device=image.device)
for b in range(B):
scale_fill = max(padded_width / float(W), padded_height / float(H)) if (W > 0 and H > 0) else 1.0
bg_w = max(1, int(round(W * scale_fill)))
bg_h = max(1, int(round(H * scale_fill)))
src_b = image[b].movedim(-1, 0).unsqueeze(0)
bg = upscale(src_b, bg_w, bg_h, "bilinear")
y0 = max(0, (bg_h - padded_height) // 2)
x0 = max(0, (bg_w - padded_width) // 2)
y1 = min(bg_h, y0 + padded_height)
x1 = min(bg_w, x0 + padded_width)
bg = bg[:, :, y0:y1, x0:x1]
if bg.shape[2] != padded_height or bg.shape[3] != padded_width:
pad_h = padded_height - bg.shape[2]
pad_w = padded_width - bg.shape[3]
pad_top_fix = max(0, pad_h // 2)
pad_bottom_fix = max(0, pad_h - pad_top_fix)
pad_left_fix = max(0, pad_w // 2)
pad_right_fix = max(0, pad_w - pad_left_fix)
bg = F.pad(bg, (pad_left_fix, pad_right_fix, pad_top_fix, pad_bottom_fix), mode="replicate")
sigma = max(1.0, 0.006 * float(min(padded_height, padded_width)))
bg = _gaussian_blur_nchw(bg, sigma_px=sigma)
if C >= 3:
r, g, bch = bg[:, 0:1], bg[:, 1:2], bg[:, 2:3]
luma = 0.2126 * r + 0.7152 * g + 0.0722 * bch
gray = torch.cat([luma, luma, luma], dim=1)
desat = 0.20
rgb = torch.cat([r, g, bch], dim=1)
rgb = rgb * (1.0 - desat) + gray * desat
bg[:, 0:3, :, :] = rgb
dim = 0.35
bg = torch.clamp(bg * dim, 0.0, 1.0)
out_image[b] = bg.squeeze(0).movedim(0, -1)
out_image[:, pad_top:pad_top+H, pad_left:pad_left+W, :] = image
# Mask handling for pillarbox_blur
if mask is not None:
fg_mask = mask
out_masks = torch.ones((B, padded_height, padded_width), dtype=image.dtype, device=image.device)
out_masks[:, pad_top:pad_top+H, pad_left:pad_left+W] = fg_mask
else:
out_masks = torch.ones((B, padded_height, padded_width), dtype=image.dtype, device=image.device)
out_masks[:, pad_top:pad_top+H, pad_left:pad_left+W] = 0.0
return (out_image, out_masks)
# Standard pad logic (edge/color)
out_image = torch.zeros((B, padded_height, padded_width, C), dtype=image.dtype, device=image.device)
for b in range(B):
if pad_mode == "edge":
# Pad with edge color
# Define edge pixels
# Pad with edge color (mean)
top_edge = image[b, 0, :, :]
bottom_edge = image[b, H-1, :, :]
left_edge = image[b, :, 0, :]
right_edge = image[b, :, W-1, :]
# Fill borders with edge colors
out_image[b, :pad_top, :, :] = top_edge.mean(dim=0)
out_image[b, pad_top+H:, :, :] = bottom_edge.mean(dim=0)
out_image[b, :, :pad_left, :] = left_edge.mean(dim=0)
out_image[b, :, pad_left+W:, :] = right_edge.mean(dim=0)
out_image[b, pad_top:pad_top+H, pad_left:pad_left+W, :] = image[b]
elif pad_mode == "edge_pixel":
# Pad with exact edge pixel values
for y in range(pad_top):
out_image[b, y, pad_left:pad_left+W, :] = image[b, 0, :, :]
for y in range(pad_top+H, padded_height):
out_image[b, y, pad_left:pad_left+W, :] = image[b, H-1, :, :]
for x in range(pad_left):
out_image[b, pad_top:pad_top+H, x, :] = image[b, :, 0, :]
for x in range(pad_left+W, padded_width):
out_image[b, pad_top:pad_top+H, x, :] = image[b, :, W-1, :]
out_image[b, :pad_top, :pad_left, :] = image[b, 0, 0, :]
out_image[b, :pad_top, pad_left+W:, :] = image[b, 0, W-1, :]
out_image[b, pad_top+H:, :pad_left, :] = image[b, H-1, 0, :]
out_image[b, pad_top+H:, pad_left+W:, :] = image[b, H-1, W-1, :]
out_image[b, pad_top:pad_top+H, pad_left:pad_left+W, :] = image[b]
else:
# Pad with specified background color
out_image[b, :, :, :] = bg_color.unsqueeze(0).unsqueeze(0) # Expand for H and W dimensions
out_image[b, :, :, :] = bg_color.unsqueeze(0).unsqueeze(0)
out_image[b, pad_top:pad_top+H, pad_left:pad_left+W, :] = image[b]
if mask is not None:
@@ -801,6 +881,10 @@ class ImagePad:
# Adapted from KJNodes, credits to Kijai
# Differences:
# - The color is an string that support various formats
# - We can copy the size of a reference image (found in V1, not in V2)
# - Fixed: input image size is copied only when both width and height are 0, allowing for one to be 0 in resize
class ImageResize:
"""
A resize and crop node, from ImageResizeKJv2
@@ -809,18 +893,30 @@ class ImageResize:
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"image": ("IMAGE", {"tooltip": "Image to resize"}),
"width": SIZE_OPT,
"height": SIZE_OPT,
"upscale_method": UPSCALE_OPT,
"keep_proportion": (["stretch", "resize", "pad", "pad_edge", "crop"], {"default": False}),
"keep_proportion": (["stretch", "resize", "pad", "pad_edge", "pad_edge_pixel", "crop", "pillarbox_blur"],
{"default": "stretch",
"tooltip": "`stretch` doesn't keep the aspect ratio\n"
"`pad` adds `pad_color` bars\n"
"`pad_edge` fills using the edge color\n"
"`resize` always keeps aspect, so W and H might change\n"
"`crop` takes a portion of the image"}),
"pad_color": COLOR_OPT,
"crop_position": (["center", "top", "bottom", "left", "right"], {"default": "center"}),
"divisible_by": ("INT", {"default": 2, "min": 0, "max": 512, "step": 1}),
"crop_position": (["center", "top", "bottom", "left", "right"],
{"default": "center", "tooltip": "Also used for `pad`"}),
"divisible_by": ("INT", {"default": 2, "min": 0, "max": 512, "step": 1,
"tooltip": "Force the final size to be divisible by"}),
},
"optional": {
"mask": ("MASK",),
"mask": ("MASK", {"tooltip": "Optional mask for the image\nwill be resized"}),
"device": (["cpu", "gpu"],),
"get_image_size": ("IMAGE", {"tooltip": "Image size to use as reference"}),
"per_batch": ("INT", {
"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 1,
"tooltip": "Process images in sub-batches to reduce memory usage. 0 disables sub-batching."}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
@@ -832,14 +928,14 @@ class ImageResize:
FUNCTION = "resize"
CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY
DESCRIPTION = ("Resizes the image to the specified width and height.\n"
"Size can be retrieved from the input.\n\n"
"Size can be retrieved from the input (when w=h=0) or a reference image.\n\n"
"Keep proportions keeps the aspect ratio of the image, by\n"
"highest dimension.")
UNIQUE_NAME = "SET_ImageResize"
DISPLAY_NAME = "Resize Image (KJ/SET)"
def resize(self, image, width, height, keep_proportion, upscale_method, divisible_by, pad_color, crop_position,
unique_id, device="cpu", mask=None):
unique_id, device="cpu", mask=None, get_image_size=None, per_batch=0):
B, H, W, C = image.shape
if device == "gpu":
@@ -849,26 +945,40 @@ class ImageResize:
else:
device = torch.device("cpu")
if width == 0:
width = W
if height == 0:
height = H
# Image size from a reference image
if get_image_size is not None:
width = get_image_size.shape[2]
height = get_image_size.shape[1]
if keep_proportion == "resize" or keep_proportion.startswith("pad"):
# Both 0 is used to copy input size
if width == 0 and height == 0:
width, height = W, H
pillarbox_blur = keep_proportion == "pillarbox_blur"
# Initialize padding variables
pad_left = pad_right = pad_top = pad_bottom = 0
# Solve the size for the ones that keeps aspect: resize, pad and pad_edge
if keep_proportion == "resize" or keep_proportion.startswith("pad") or pillarbox_blur:
# If one of the dimensions is zero, calculate it to maintain the aspect ratio
if width == 0 and height != 0:
ratio = height / H
new_width = round(W * ratio)
new_height = height
elif height == 0 and width != 0:
ratio = width / W
new_width = width
new_height = round(H * ratio)
elif width != 0 and height != 0:
# Scale based on which dimension is smaller in proportion to the desired dimensions
ratio = min(width / W, height / H)
new_width = round(W * ratio)
new_height = round(H * ratio)
else:
new_width = width
new_height = height
if keep_proportion.startswith("pad"):
if keep_proportion.startswith("pad") or pillarbox_blur:
# Calculate padding based on position
if crop_position == "center":
pad_left = (width - new_width) // 2
@@ -903,60 +1013,77 @@ class ImageResize:
width = width - (width % divisible_by)
height = height - (height % divisible_by)
out_image = image.clone().to(device)
# Preflight estimate (log-only when batching is active)
if per_batch != 0 and B > per_batch:
try:
bytes_per_elem = image.element_size() # typically 4 for float32
est_total_bytes = B * height * width * C * bytes_per_elem
est_mb = est_total_bytes / (1024 * 1024)
msg = f"<tr><td>Resize Image</td><td>estimated output ~{est_mb:.2f} MB; batching {per_batch}/{B}</td></tr>"
if unique_id and PromptServer is not None:
try:
PromptServer.instance.send_progress_text(msg, unique_id)
except Exception:
pass
logger.info(f"estimated output ~{est_mb:.2f} MB; batching {per_batch}/{B}")
except Exception:
pass
if mask is not None:
out_mask = mask.clone().to(device)
else:
out_mask = None
def _process_subbatch(in_image, in_mask, pad_left, pad_right, pad_top, pad_bottom):
# Avoid unnecessary clones; only move if needed
out_image = in_image if in_image.device == device else in_image.to(device)
out_mask = None if in_mask is None else (in_mask if in_mask.device == device else in_mask.to(device))
if keep_proportion == "crop":
old_width = W
old_height = H
old_aspect = old_width / old_height
new_aspect = width / height
# Crop logic
if keep_proportion == "crop":
old_height = out_image.shape[-3]
old_width = out_image.shape[-2]
old_aspect = old_width / old_height
new_aspect = width / height
# Calculate dimensions to keep
if old_aspect > new_aspect: # Image is wider than target
crop_w = round(old_height * new_aspect)
crop_h = old_height
else: # Image is taller than target
crop_w = old_width
crop_h = round(old_width / new_aspect)
# Calculate dimensions to keep
if old_aspect > new_aspect: # Image is wider than target
crop_w = round(old_height * new_aspect)
crop_h = old_height
else: # Image is taller than target
crop_w = old_width
crop_h = round(old_width / new_aspect)
# Calculate crop position
if crop_position == "center":
x = (old_width - crop_w) // 2
y = (old_height - crop_h) // 2
elif crop_position == "top":
x = (old_width - crop_w) // 2
y = 0
elif crop_position == "bottom":
x = (old_width - crop_w) // 2
y = old_height - crop_h
elif crop_position == "left":
x = 0
y = (old_height - crop_h) // 2
elif crop_position == "right":
x = old_width - crop_w
y = (old_height - crop_h) // 2
# Calculate crop position
if crop_position == "center":
x = (old_width - crop_w) // 2
y = (old_height - crop_h) // 2
elif crop_position == "top":
x = (old_width - crop_w) // 2
y = 0
elif crop_position == "bottom":
x = (old_width - crop_w) // 2
y = old_height - crop_h
elif crop_position == "left":
x = 0
y = (old_height - crop_h) // 2
elif crop_position == "right":
x = old_width - crop_w
y = (old_height - crop_h) // 2
# Apply crop
out_image = out_image.narrow(-2, x, crop_w).narrow(-3, y, crop_h)
if mask is not None:
out_mask = out_mask.narrow(-1, x, crop_w).narrow(-2, y, crop_h)
# Apply crop
out_image = out_image.narrow(-2, x, crop_w).narrow(-3, y, crop_h)
if out_mask is not None:
out_mask = out_mask.narrow(-1, x, crop_w).narrow(-2, y, crop_h)
out_image = upscale(out_image.movedim(-1, 1), width, height, upscale_method).movedim(1, -1)
# Resize the image
out_image = upscale(out_image.movedim(-1, 1), width, height, upscale_method).movedim(1, -1)
if mask is not None:
# if upscale_method == "lanczos":
# out_mask = upscale(out_mask.unsqueeze(1).repeat(1, 3, 1, 1), width, height,
# upscale_method).movedim(1, -1)[:, :, :, 0]
# else:
out_mask = upscale(out_mask.unsqueeze(1), width, height, upscale_method).squeeze(1)
if out_mask is not None:
if upscale_method == "lanczos":
out_mask = upscale(out_mask.unsqueeze(1).repeat(1, 3, 1, 1), width, height,
upscale_method).movedim(1, -1)[:, :, :, 0]
else:
out_mask = upscale(out_mask.unsqueeze(1), width, height, upscale_method).squeeze(1)
if keep_proportion.startswith("pad"):
if pad_left > 0 or pad_right > 0 or pad_top > 0 or pad_bottom > 0:
# Pad logic
if (keep_proportion.startswith("pad") or pillarbox_blur) and (pad_left > 0 or pad_right > 0 or pad_top > 0
or pad_bottom > 0):
padded_width = width + pad_left + pad_right
padded_height = height + pad_top + pad_bottom
if divisible_by > 1:
@@ -968,18 +1095,57 @@ class ImageResize:
if height_remainder > 0:
extra_height = divisible_by - height_remainder
pad_bottom += extra_height
out_image, _ = ImagePad.pad(self, out_image, pad_left, pad_right, pad_top, pad_bottom, 0, pad_color,
"edge" if keep_proportion == "pad_edge" else "color")
if mask is not None:
out_mask = out_mask.unsqueeze(1).repeat(1, 3, 1, 1).movedim(1, -1)
out_mask, _ = ImagePad.pad(self, out_mask, pad_left, pad_right, pad_top, pad_bottom, 0, pad_color,
"edge" if keep_proportion == "pad_edge" else "color")
out_mask = out_mask[:, :, :, 0]
else:
B, H_pad, W_pad, _ = out_image.shape
out_mask = torch.ones((B, H_pad, W_pad), dtype=out_image.dtype, device=out_image.device)
out_mask[:, pad_top:pad_top+height, pad_left:pad_left+width] = 0.0
pad_mode = (
"pillarbox_blur" if pillarbox_blur else
"edge" if keep_proportion == "pad_edge" else
"edge_pixel" if keep_proportion == "pad_edge_pixel" else
"color"
)
out_image, out_mask = ImagePad.pad(self, out_image, pad_left, pad_right, pad_top, pad_bottom, 0, pad_color,
pad_mode, mask=out_mask)
return out_image, out_mask
# If batching disabled (per_batch==0) or batch fits, process whole batch
if per_batch == 0 or B <= per_batch:
out_image, out_mask = _process_subbatch(image, mask, pad_left, pad_right, pad_top, pad_bottom)
else:
chunks = []
mask_chunks = [] if mask is not None else None
total_batches = (B + per_batch - 1) // per_batch
current_batch = 0
for start_idx in range(0, B, per_batch):
current_batch += 1
end_idx = min(start_idx + per_batch, B)
sub_img = image[start_idx:end_idx]
sub_mask = mask[start_idx:end_idx] if mask is not None else None
sub_out_img, sub_out_mask = _process_subbatch(sub_img, sub_mask, pad_left, pad_right, pad_top, pad_bottom)
chunks.append(sub_out_img.cpu())
if mask is not None:
mask_chunks.append(sub_out_mask.cpu() if sub_out_mask is not None else None)
# Per-batch progress update
if unique_id and PromptServer is not None:
try:
PromptServer.instance.send_progress_text(
f"<tr><td>Resize Image</td><td>batch {current_batch}/{total_batches} · images {end_idx}/{B}"
"</td></tr>",
unique_id
)
except Exception:
pass
else:
try:
logger.info(f"batch {current_batch}/{total_batches} · images {end_idx}/{B}")
except Exception:
pass
out_image = torch.cat(chunks, dim=0)
if mask is not None and any(m is not None for m in mask_chunks):
out_mask = torch.cat([m for m in mask_chunks if m is not None], dim=0)
else:
out_mask = None
# Progress UI
if unique_id and PromptServer is not None:
try:
num_elements = out_image.numel()
@@ -997,3 +1163,53 @@ class ImageResize:
return (out_image.cpu(), out_image.shape[2], out_image.shape[1],
out_mask.cpu() if out_mask is not None else
torch.zeros(64, 64, device=torch.device("cpu"), dtype=torch.float32))
# Adapted from KJNodes, credits to Kijai
# Difference: reference image `get_image_size`
class ResizeMask:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask": ("MASK",),
"width": SIZE_OPT,
"height": SIZE_OPT,
"keep_proportions": ("BOOLEAN", {"default": False}),
"upscale_method": UPSCALE_OPT_MASK,
"crop": (["disabled", "center"],),
},
"optional": {
"get_image_size": ("IMAGE", {"tooltip": "Image size to use as reference"}),
},
}
RETURN_TYPES = ("MASK", "INT", "INT",)
RETURN_NAMES = ("mask", "width", "height",)
FUNCTION = "resize"
CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY
DESCRIPTION = "Resizes the mask or batch of masks to the specified width and height."
UNIQUE_NAME = "SET_ResizeMask"
DISPLAY_NAME = "Resize Mask (KJ/SET)"
def resize(self, mask, width, height, keep_proportions, upscale_method, crop, get_image_size=None):
# Image size from a reference image
if get_image_size is not None:
width = get_image_size.shape[2]
height = get_image_size.shape[1]
if keep_proportions:
_, oh, ow = mask.shape
width = ow if width == 0 else width
height = oh if height == 0 else height
ratio = min(width / ow, height / oh)
width = round(ow*ratio)
height = round(oh*ratio)
if upscale_method == "lanczos":
out_mask = common_upscale(mask.unsqueeze(1).repeat(1, 3, 1, 1), width, height, upscale_method,
crop=crop).movedim(1, -1)[:, :, :, 0]
else:
out_mask = common_upscale(mask.unsqueeze(1), width, height, upscale_method, crop=crop).squeeze(1)
return (out_mask, out_mask.shape[2], out_mask.shape[1],)