diff --git a/src/nodes/nodes_img.py b/src/nodes/nodes_img.py index 3bea559..a028fa3 100644 --- a/src/nodes/nodes_img.py +++ b/src/nodes/nodes_img.py @@ -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"Resize Imageestimated output ~{est_mb:.2f} MB; batching {per_batch}/{B}" + 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"Resize Imagebatch {current_batch}/{total_batches} · images {end_idx}/{B}" + "", + 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],)