diff --git a/src/nodes/nodes_img.py b/src/nodes/nodes_img.py index fe6cdd8..3bea559 100644 --- a/src/nodes/nodes_img.py +++ b/src/nodes/nodes_img.py @@ -2,6 +2,7 @@ # 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 import numpy as np @@ -12,18 +13,36 @@ from seconohe.foreground_estimation.affce import affce from seconohe.foreground_estimation.fmlfe import fmlfe, IMPL_PRIORITY from seconohe.color import color_to_rgb_float from seconohe.downloader import download_file +from seconohe.color import color_to_rgb_uint8 # We are the main source, so we use the main_logger 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: from folder_paths import get_input_directory # To get the ComfyUI input directory from comfy import model_management + from comfy.utils import common_upscale except ModuleNotFoundError: # No ComfyUI, this is a test environment def get_input_directory(): return "" + +try: + from nodes import ImageScale +except Exception: + class ImageScale(object): + upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"] +try: + from nodes import MAX_RESOLUTION +except Exception: + MAX_RESOLUTION = 16384 +try: + from server import PromptServer +except ModuleNotFoundError: + PromptServer = None try: # We need to import the built-in LoadImage class for ImageDownload from nodes import LoadImage @@ -44,6 +63,17 @@ 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": DEFAULT_UPSCALE, + "tooltip": "Interpolation method for image resize" + }) +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_FI[1]["forceInput"] = True +MASK_UPSCALE = 'nearest-exact' # transforms.InterpolationMode.NEAREST_EXACT.value +BEST_UPSCALE = 'lanczos' # transforms.InterpolationMode.LANCZOS.value def tensor_to_pil(tensor: torch.Tensor) -> Image.Image: @@ -58,6 +88,11 @@ def pil_to_tensor(pil_image: Image.Image) -> torch.Tensor: return torch.from_numpy(np_image) +def upscale(image, width, height, upscale_method): + # return F.interpolate(image, size=(height, width), mode=upscale_method) + return common_upscale(image, width, height, upscale_method, crop="disabled") + + if has_load_image: class ImageDownload: @classmethod @@ -663,3 +698,302 @@ class CreateEmptyImage: final_image = color_tensor.expand(b, h, w, 3) return (final_image,) + + +# Adapted from KJNodes, credits to Kijai +class ImagePad: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE", ), + "left": PAD_SIZE_OPT, + "right": PAD_SIZE_OPT, + "top": PAD_SIZE_OPT, + "bottom": PAD_SIZE_OPT, + "extra_padding": PAD_SIZE_OPT, + "pad_mode": (["edge", "color"],), + "color": COLOR_OPT, + }, + "optional": { + "mask": ("MASK", ), + "target_width": SIZE_OPT_FI, + "target_height": SIZE_OPT_FI, + } + } + + RETURN_TYPES = ("IMAGE", "MASK", ) + RETURN_NAMES = ("images", "masks",) + FUNCTION = "pad" + CATEGORY = BASE_CATEGORY + "/" + MANIPULATION_CATEGORY + DESCRIPTION = "Pad the input image and optionally mask with the specified padding." + UNIQUE_NAME = "SET_ImagePad" + DISPLAY_NAME = "Pad Image (KJ/SET)" + + def pad(self, image, left, right, top, bottom, extra_padding, color, pad_mode, mask=None, target_width=None, + target_height=None): + B, H, W, C = image.shape + + # Resize masks to image dimensions if necessary + if mask is not None: + BM, HM, WM = mask.shape + if HM != H or WM != W: + mask = F.interpolate(mask.unsqueeze(1), size=(H, W), mode=MASK_UPSCALE).squeeze(1) + + # Parse background color + bg_color = torch.tensor(color_to_rgb_uint8(logger, color), dtype=image.dtype, device=image.device) + + # Calculate padding sizes with extra padding + if target_width is not None and target_height is not None: + if extra_padding > 0: + image = upscale(image.movedim(-1, 1), W - extra_padding, H - extra_padding, BEST_UPSCALE).movedim(1, -1) + B, H, W, C = image.shape + + padded_width = target_width + padded_height = target_height + pad_left = (padded_width - W) // 2 + pad_right = padded_width - W - pad_left + pad_top = (padded_height - H) // 2 + pad_bottom = padded_height - H - pad_top + else: + pad_left = left + extra_padding + pad_right = right + extra_padding + pad_top = top + extra_padding + pad_bottom = bottom + extra_padding + + 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 + for b in range(B): + if pad_mode == "edge": + # Pad with edge color + # Define edge pixels + 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] + else: + # Pad with specified background color + out_image[b, :, :, :] = bg_color.unsqueeze(0).unsqueeze(0) # Expand for H and W dimensions + out_image[b, pad_top:pad_top+H, pad_left:pad_left+W, :] = image[b] + + if mask is not None: + out_masks = torch.nn.functional.pad( + mask, + (pad_left, pad_right, pad_top, pad_bottom), + mode='replicate' + ) + else: + out_masks = torch.ones((B, padded_height, padded_width), dtype=image.dtype, device=image.device) + for m in range(B): + out_masks[m, pad_top:pad_top+H, pad_left:pad_left+W] = 0.0 + + return (out_image, out_masks) + + +# Adapted from KJNodes, credits to Kijai +class ImageResize: + """ + A resize and crop node, from ImageResizeKJv2 + """ + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "width": SIZE_OPT, + "height": SIZE_OPT, + "upscale_method": UPSCALE_OPT, + "keep_proportion": (["stretch", "resize", "pad", "pad_edge", "crop"], {"default": False}), + "pad_color": COLOR_OPT, + "crop_position": (["center", "top", "bottom", "left", "right"], {"default": "center"}), + "divisible_by": ("INT", {"default": 2, "min": 0, "max": 512, "step": 1}), + }, + "optional": { + "mask": ("MASK",), + "device": (["cpu", "gpu"],), + }, + "hidden": { + "unique_id": "UNIQUE_ID", + }, + } + + RETURN_TYPES = ("IMAGE", "INT", "INT", "MASK",) + RETURN_NAMES = ("IMAGE", "width", "height", "mask",) + 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" + "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): + B, H, W, C = image.shape + + if device == "gpu": + if upscale_method == "lanczos": + raise Exception("Lanczos is not supported on the GPU") + device = model_management.get_torch_device() + else: + device = torch.device("cpu") + + if width == 0: + width = W + if height == 0: + height = H + + if keep_proportion == "resize" or keep_proportion.startswith("pad"): + # 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) + elif height == 0 and width != 0: + ratio = width / W + 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) + + if keep_proportion.startswith("pad"): + # Calculate padding based on position + if crop_position == "center": + pad_left = (width - new_width) // 2 + pad_right = width - new_width - pad_left + pad_top = (height - new_height) // 2 + pad_bottom = height - new_height - pad_top + elif crop_position == "top": + pad_left = (width - new_width) // 2 + pad_right = width - new_width - pad_left + pad_top = 0 + pad_bottom = height - new_height + elif crop_position == "bottom": + pad_left = (width - new_width) // 2 + pad_right = width - new_width - pad_left + pad_top = height - new_height + pad_bottom = 0 + elif crop_position == "left": + pad_left = 0 + pad_right = width - new_width + pad_top = (height - new_height) // 2 + pad_bottom = height - new_height - pad_top + elif crop_position == "right": + pad_left = width - new_width + pad_right = 0 + pad_top = (height - new_height) // 2 + pad_bottom = height - new_height - pad_top + + width = new_width + height = new_height + + if divisible_by > 1: + width = width - (width % divisible_by) + height = height - (height % divisible_by) + + out_image = image.clone().to(device) + + if mask is not None: + out_mask = mask.clone().to(device) + else: + out_mask = None + + if keep_proportion == "crop": + old_width = W + old_height = H + 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 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) + + 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 keep_proportion.startswith("pad"): + if 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: + width_remainder = padded_width % divisible_by + height_remainder = padded_height % divisible_by + if width_remainder > 0: + extra_width = divisible_by - width_remainder + pad_right += extra_width + 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 + + if unique_id and PromptServer is not None: + try: + num_elements = out_image.numel() + element_size = out_image.element_size() + memory_size_mb = (num_elements * element_size) / (1024 * 1024) + + PromptServer.instance.send_progress_text( + f"