From 12349affd602e8ca5902c783fa133b43a804689b Mon Sep 17 00:00:00 2001 From: Cyber Dick Lang <286878701@qq.com> Date: Wed, 2 Jul 2025 20:47:56 +0800 Subject: [PATCH] Add gap parameter to image concatenate nodes Introduces a 'gap' parameter to both image_concatenate.py and image_concatenate_multi.py, allowing users to specify spacing between concatenated images. Also updates the imitation_hue_node.py to process batches with progress bar support and corrects mask handling for multiple target images. --- nodes/image/image_concatenate.py | 33 +- nodes/image/image_concatenate_multi.py | 399 +++++++++++++++---------- nodes/image/imitation_hue_node.py | 44 +-- 3 files changed, 292 insertions(+), 184 deletions(-) diff --git a/nodes/image/image_concatenate.py b/nodes/image/image_concatenate.py index 4a1cfde..cc29b31 100644 --- a/nodes/image/image_concatenate.py +++ b/nodes/image/image_concatenate.py @@ -32,6 +32,7 @@ class ImageConcatenate_UTK: }), "match_image_size": ("BOOLEAN", {"default": True}), "max_size": ("INT", {"default": 4096, "min": 64, "max": 8192, "step": 64}), + "gap": ("INT", {"default": 0, "min": 0, "max": 512, "step": 1}), "background_color": (["black", "white", "gray", "transparent"], {"default": "black"}), } } @@ -39,7 +40,7 @@ class ImageConcatenate_UTK: RETURN_TYPES = ("IMAGE",) FUNCTION = "concatenate" - def concatenate(self, image1, image2, direction, match_image_size, max_size, background_color): + def concatenate(self, image1, image2, direction, match_image_size, max_size, gap, background_color): # Check if the batch sizes are different batch_size1 = image1.shape[0] batch_size2 = image2.shape[0] @@ -111,12 +112,12 @@ class ImageConcatenate_UTK: h1, w1 = image1.shape[1:3] h2, w2 = image2.shape[1:3] - # Calculate final dimensions + # Calculate final dimensions with gap if direction in ['right', 'left']: final_height = max(h1, h2) - final_width = w1 + w2 + final_width = w1 + w2 + (gap if gap > 0 else 0) else: # up, down - final_height = h1 + h2 + final_height = h1 + h2 + (gap if gap > 0 else 0) final_width = max(w1, w2) # Check if we need to scale down @@ -142,9 +143,9 @@ class ImageConcatenate_UTK: h2, w2 = image2.shape[1:3] if direction in ['right', 'left']: final_height = max(h1, h2) - final_width = w1 + w2 + final_width = w1 + w2 + (gap if gap > 0 else 0) else: - final_height = h1 + h2 + final_height = h1 + h2 + (gap if gap > 0 else 0) final_width = max(w1, w2) # Ensure both images have the same number of channels @@ -160,20 +161,24 @@ class ImageConcatenate_UTK: # 创建输出张量,batch维度与输入一致 batch_size = image1.shape[0] - if background_color == "transparent": - output = torch.zeros((batch_size, final_height, final_width, image1.shape[-1]), dtype=image1.dtype, device=image1.device) + if gap > 0: + if background_color == "transparent": + output = torch.zeros((batch_size, final_height, final_width, image1.shape[-1]), dtype=image1.dtype, device=image1.device) + else: + color_value = 1.0 if background_color == "white" else 0.0 if background_color == "black" else 0.5 + output = torch.full((batch_size, final_height, final_width, image1.shape[-1]), color_value, dtype=image1.dtype, device=image1.device) else: - color_value = 1.0 if background_color == "white" else 0.0 if background_color == "black" else 0.5 - output = torch.full((batch_size, final_height, final_width, image1.shape[-1]), color_value, dtype=image1.dtype, device=image1.device) + # gap=0时,保持原有逻辑,默认黑色背景 + output = torch.zeros((batch_size, final_height, final_width, image1.shape[-1]), dtype=image1.dtype, device=image1.device) # 计算放置位置 if direction == 'right': x1 = 0 - x2 = w1 + x2 = w1 + (gap if gap > 0 else 0) y1 = (final_height - h1) // 2 y2 = (final_height - h2) // 2 elif direction == 'left': - x1 = w2 + x1 = w2 + (gap if gap > 0 else 0) x2 = 0 y1 = (final_height - h1) // 2 y2 = (final_height - h2) // 2 @@ -181,11 +186,11 @@ class ImageConcatenate_UTK: x1 = (final_width - w1) // 2 x2 = (final_width - w2) // 2 y1 = 0 - y2 = h1 + y2 = h1 + (gap if gap > 0 else 0) else: # up x1 = (final_width - w1) // 2 x2 = (final_width - w2) // 2 - y1 = h2 + y1 = h2 + (gap if gap > 0 else 0) y2 = 0 # 批量放置图片 diff --git a/nodes/image/image_concatenate_multi.py b/nodes/image/image_concatenate_multi.py index 0a77672..83ee114 100644 --- a/nodes/image/image_concatenate_multi.py +++ b/nodes/image/image_concatenate_multi.py @@ -1,195 +1,288 @@ """ -Image Concatenate Multi Node +Image Concatenate Multi Node (UTK) ~~~~~~~~~~~~~~~~~~~~~~~~~~~ -Concatenates multiple images in various directions and layouts. +版本: 1.3.0 +最后更新: 2024-06-09 +作者: May + +变更日志: +- v1.3.0: 移除grid_size参数,增加gap参数,支持拼接间距设置。 +- v1.2.0: 默认启用智能拼接(smart),支持2~4图,自动等比缩放,逐步拼接,输出接近正方形。 +- v1.1.x: 支持sequential/smart两种模式,拼接逻辑优化。 +- v1.0.0: 初始版本,基础多图拼接。 :copyright: (c) 2024 by May :license: MIT, see LICENSE for more details. """ import torch +import math class ImageConcatenateMulti_UTK: - CATEGORY = "UniversalToolkit/Image" - @classmethod def INPUT_TYPES(cls): return { "required": { - "images": ("IMAGE",), + "image_1": ("IMAGE", ), + "image_2": ("IMAGE", ), + "mode": (["sequential", "smart"], {"default": "smart"}), "direction": ( - [ 'right', - 'down', - 'left', - 'up', - 'auto', - ], - { - "default": 'auto' - }), - "match_image_size": ("BOOLEAN", {"default": True}), + ['right', 'down', 'left', 'up'], + {"default": 'right'} + ), + "match_image_size": ("BOOLEAN", {"default": False}), "max_size": ("INT", {"default": 4096, "min": 64, "max": 8192, "step": 64}), "background_color": (["black", "white", "gray", "transparent"], {"default": "black"}), - "grid_size": (["auto", "1x1", "2x2", "3x3", "4x4"], {"default": "auto"}), + "gap": ("INT", {"default": 0, "min": 0, "max": 256, "step": 1}), + }, + "optional": { + "image_3": ("IMAGE", ), + "image_4": ("IMAGE", ), } } - + RETURN_TYPES = ("IMAGE",) - FUNCTION = "concatenate_multi" - - def concatenate_multi(self, images, direction, match_image_size, max_size, background_color, grid_size): - if len(images.shape) != 4: - raise ValueError("输入必须是4D张量 [batch, height, width, channels]") - - batch_size = images.shape[0] - - # 处理网格布局 - if grid_size != "auto": - rows, cols = map(int, grid_size.split("x")) - if batch_size > rows * cols: - raise ValueError(f"图像数量 ({batch_size}) 超过网格大小 ({grid_size})") - # 填充到网格大小 - if batch_size < rows * cols: - padding = torch.zeros((rows * cols - batch_size, *images.shape[1:]), dtype=images.dtype, device=images.device) - images = torch.cat([images, padding], dim=0) - batch_size = rows * cols - else: - # 自动计算网格大小 - if direction in ['right', 'left', 'auto']: - rows = 1 - cols = batch_size - else: # up, down - rows = batch_size - cols = 1 + RETURN_NAMES = ("images",) + FUNCTION = "combine" + CATEGORY = "KJNodes/image" + DESCRIPTION = """ +Creates an image from multiple images. +你可以设置inputcount来决定输入端口数量, +并支持顺序拼接(sequential)和正方形智能布局(square)两种模式。 +""" - # 获取所有图像的尺寸 - heights = [] - widths = [] - for i in range(batch_size): - h, w = images[i].shape[:2] - heights.append(h) - widths.append(w) + def combine(self, mode, direction, match_image_size, max_size, background_color, gap, image_1, image_2, image_3=None, image_4=None): + images = [image_1, image_2] + if image_3 is not None: + images.append(image_3) + if image_4 is not None: + images.append(image_4) + if mode == "smart": + img = images[0] + for ni in images[1:]: + img, = self._smart_pair_concatenate(img, ni, max_size, background_color, gap) + return (img,) + image = images[0] + for new_image in images[1:]: + image, = self._sequential_concatenate( + image, new_image, direction, match_image_size, max_size, background_color, gap + ) + return (image,) - # 如果方向是auto,计算最佳方向 + def _sequential_concatenate(self, image1, image2, direction, match_image_size, max_size, background_color, gap): + h1, w1 = image1.shape[1:3] + h2, w2 = image2.shape[1:3] if direction == 'auto': - # 计算水平和垂直拼接的宽高比 - total_width = sum(widths) - max_height = max(heights) - horizontal_ratio = total_width / max_height - - total_height = sum(heights) - max_width = max(widths) - vertical_ratio = max_width / total_height - - # 选择更接近1:1的方向 + horizontal_ratio = (w1 + w2) / max(h1, h2) + vertical_ratio = max(w1, w2) / (h1 + h2) direction = 'right' if abs(horizontal_ratio - 1) <= abs(vertical_ratio - 1) else 'down' - - # 如果需要匹配图像尺寸 if match_image_size: - if direction in ['right', 'left', 'auto']: - # 匹配高度 - target_height = max(heights) - for i in range(batch_size): - if heights[i] < target_height: - scale = target_height / heights[i] - new_width = int(widths[i] * scale) - images[i] = torch.nn.functional.interpolate( - images[i].unsqueeze(0).permute(0, 3, 1, 2), - size=(target_height, new_width), - mode='bilinear', - align_corners=False - ).permute(0, 2, 3, 1).squeeze(0) - else: # up, down - # 匹配宽度 - target_width = max(widths) - for i in range(batch_size): - if widths[i] < target_width: - scale = target_width / widths[i] - new_height = int(heights[i] * scale) - images[i] = torch.nn.functional.interpolate( - images[i].unsqueeze(0).permute(0, 3, 1, 2), - size=(new_height, target_width), - mode='bilinear', - align_corners=False - ).permute(0, 2, 3, 1).squeeze(0) - - # 更新尺寸 - heights = [] - widths = [] - for i in range(batch_size): - h, w = images[i].shape[:2] - heights.append(h) - widths.append(w) - - # 计算最终输出尺寸 - if direction in ['right', 'left', 'auto']: - final_height = max(heights) - final_width = sum(widths) - else: # up, down - final_height = sum(heights) - final_width = max(widths) - - # 检查是否需要缩放 + if direction in ['right', 'left']: + target_height = max(h1, h2) + if h1 < target_height: + scale = target_height / h1 + new_width = int(w1 * scale) + image1 = torch.nn.functional.interpolate( + image1.permute(0, 3, 1, 2), + size=(target_height, new_width), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1) + if h2 < target_height: + scale = target_height / h2 + new_width = int(w2 * scale) + image2 = torch.nn.functional.interpolate( + image2.permute(0, 3, 1, 2), + size=(target_height, new_width), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1) + else: + target_width = max(w1, w2) + if w1 < target_width: + scale = target_width / w1 + new_height = int(h1 * scale) + image1 = torch.nn.functional.interpolate( + image1.permute(0, 3, 1, 2), + size=(new_height, target_width), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1) + if w2 < target_width: + scale = target_width / w2 + new_height = int(h2 * scale) + image2 = torch.nn.functional.interpolate( + image2.permute(0, 3, 1, 2), + size=(new_height, target_width), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1) + h1, w1 = image1.shape[1:3] + h2, w2 = image2.shape[1:3] + if direction in ['right', 'left']: + final_height = max(h1, h2) + final_width = w1 + w2 + (gap if gap > 0 else 0) + else: + final_height = h1 + h2 + (gap if gap > 0 else 0) + final_width = max(w1, w2) if max(final_height, final_width) > max_size: scale = max_size / max(final_height, final_width) - for i in range(batch_size): - new_height = int(heights[i] * scale) - new_width = int(widths[i] * scale) - images[i] = torch.nn.functional.interpolate( - images[i].unsqueeze(0).permute(0, 3, 1, 2), - size=(new_height, new_width), - mode='bilinear', - align_corners=False - ).permute(0, 2, 3, 1).squeeze(0) + new_h1 = int(h1 * scale) + new_w1 = int(w1 * scale) + new_h2 = int(h2 * scale) + new_w2 = int(w2 * scale) + image1 = torch.nn.functional.interpolate( + image1.permute(0, 3, 1, 2), + size=(new_h1, new_w1), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1) + image2 = torch.nn.functional.interpolate( + image2.permute(0, 3, 1, 2), + size=(new_h2, new_w2), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1) + h1, w1 = image1.shape[1:3] + h2, w2 = image2.shape[1:3] + if direction in ['right', 'left']: + final_height = max(h1, h2) + final_width = w1 + w2 + (gap if gap > 0 else 0) + else: + final_height = h1 + h2 + (gap if gap > 0 else 0) + final_width = max(w1, w2) + if background_color == "transparent": + output = torch.zeros((1, final_height, final_width, image1.shape[-1]), dtype=image1.dtype, device=image1.device) + else: + color_value = 1.0 if background_color == "white" else 0.0 if background_color == "black" else 0.5 + output = torch.full((1, final_height, final_width, image1.shape[-1]), color_value, dtype=image1.dtype, device=image1.device) + if direction == 'right': + x1 = 0 + x2 = w1 + (gap if gap > 0 else 0) + y1 = (final_height - h1) // 2 + y2 = (final_height - h2) // 2 + elif direction == 'left': + x1 = w2 + (gap if gap > 0 else 0) + x2 = 0 + y1 = (final_height - h1) // 2 + y2 = (final_height - h2) // 2 + elif direction == 'down': + x1 = (final_width - w1) // 2 + x2 = (final_width - w2) // 2 + y1 = 0 + y2 = h1 + (gap if gap > 0 else 0) + else: # up + x1 = (final_width - w1) // 2 + x2 = (final_width - w2) // 2 + y1 = h2 + (gap if gap > 0 else 0) + y2 = 0 + output[:, y1:y1+h1, x1:x1+w1] = image1 + output[:, y2:y2+h2, x2:x2+w2] = image2 + return (output,) - # 更新最终尺寸 - heights = [] - widths = [] + def _square_concatenate(self, images, max_size, background_color): + batch_size = images.shape[0] + grid_size = math.ceil(math.sqrt(batch_size)) + rows = cols = grid_size + original_heights = [] + original_widths = [] for i in range(batch_size): h, w = images[i].shape[:2] - heights.append(h) - widths.append(w) - - if direction in ['right', 'left', 'auto']: - final_height = max(heights) - final_width = sum(widths) - else: # up, down - final_height = sum(heights) - final_width = max(widths) - - # 创建输出张量 + original_heights.append(h) + original_widths.append(w) + total_area = sum(h * w for h, w in zip(original_heights, original_widths)) + target_canvas_size = math.sqrt(total_area) + if target_canvas_size > max_size: + target_canvas_size = max_size + cell_width = int(target_canvas_size / cols) + cell_height = int(target_canvas_size / rows) + cell_width = max(cell_width, 64) + cell_height = max(cell_height, 64) + final_width = cell_width * cols + final_height = cell_height * rows if background_color == "transparent": output = torch.zeros((1, final_height, final_width, images.shape[-1]), dtype=images.dtype, device=images.device) else: color_value = 1.0 if background_color == "white" else 0.0 if background_color == "black" else 0.5 output = torch.full((1, final_height, final_width, images.shape[-1]), color_value, dtype=images.dtype, device=images.device) + for i in range(batch_size): + row = i // cols + col = i % cols + x_offset = col * cell_width + y_offset = row * cell_height + scaled_image = torch.nn.functional.interpolate( + images[i].unsqueeze(0).permute(0, 3, 1, 2), + size=(cell_height, cell_width), + mode='bilinear', + align_corners=False + ).permute(0, 2, 3, 1).squeeze(0) + output[0, y_offset:y_offset+cell_height, x_offset:x_offset+cell_width] = scaled_image + return (output,) - # 放置图像 - if direction in ['right', 'left', 'auto']: - x_offset = 0 - for i in range(batch_size): - h, w = images[i].shape[:2] - y_offset = (final_height - h) // 2 - if direction == 'left': - x_offset = final_width - sum(widths[i:]) - output[0, y_offset:y_offset+h, x_offset:x_offset+w] = images[i] - if direction != 'left': - x_offset += w - else: # up, down - y_offset = 0 - for i in range(batch_size): - h, w = images[i].shape[:2] - x_offset = (final_width - w) // 2 - if direction == 'up': - y_offset = final_height - sum(heights[i:]) - output[0, y_offset:y_offset+h, x_offset:x_offset+w] = images[i] - if direction != 'up': - y_offset += h - + def _smart_pair_concatenate(self, img1, img2, max_size, background_color, gap): + b, h1, w1, c = img1.shape + b2, h2, w2, c2 = img2.shape + target_height = max(h1, h2) + scale1_h = target_height / h1 + scale2_h = target_height / h2 + w1_h = int(w1 * scale1_h) + w2_h = int(w2 * scale2_h) + target_width = max(w1, w2) + scale1_w = target_width / w1 + scale2_w = target_width / w2 + h1_w = int(h1 * scale1_w) + h2_w = int(h2 * scale2_w) + out_w_h = w1_h + w2_h + (gap if gap > 0 else 0) + out_h_h = target_height + out_w_w = target_width + out_h_w = h1_w + h2_w + (gap if gap > 0 else 0) + ratio_h = max(out_w_h, out_h_h) / min(out_w_h, out_h_h) + ratio_w = max(out_w_w, out_h_w) / min(out_w_w, out_h_w) + if abs(ratio_h - 1) <= abs(ratio_w - 1): + img1r = torch.nn.functional.interpolate(img1.permute(0, 3, 1, 2), size=(target_height, w1_h), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + img2r = torch.nn.functional.interpolate(img2.permute(0, 3, 1, 2), size=(target_height, w2_h), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + final_height = target_height + final_width = w1_h + w2_h + (gap if gap > 0 else 0) + if max(final_height, final_width) > max_size: + scale = max_size / max(final_height, final_width) + new_height = int(final_height * scale) + new_w1 = int(w1_h * scale) + new_w2 = int(w2_h * scale) + img1r = torch.nn.functional.interpolate(img1r.permute(0, 3, 1, 2), size=(new_height, new_w1), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + img2r = torch.nn.functional.interpolate(img2r.permute(0, 3, 1, 2), size=(new_height, new_w2), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + final_height = new_height + final_width = new_w1 + new_w2 + (gap if gap > 0 else 0) + if background_color == "transparent": + output = torch.zeros((1, final_height, final_width, img1.shape[-1]), dtype=img1.dtype, device=img1.device) + else: + color_value = 1.0 if background_color == "white" else 0.0 if background_color == "black" else 0.5 + output = torch.full((1, final_height, final_width, img1.shape[-1]), color_value, dtype=img1.dtype, device=img1.device) + output[:, :, :img1r.shape[2]] = img1r + output[:, :, img1r.shape[2]+(gap if gap > 0 else 0):] = img2r + else: + img1r = torch.nn.functional.interpolate(img1.permute(0, 3, 1, 2), size=(h1_w, target_width), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + img2r = torch.nn.functional.interpolate(img2.permute(0, 3, 1, 2), size=(h2_w, target_width), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + final_height = h1_w + h2_w + (gap if gap > 0 else 0) + final_width = target_width + if max(final_height, final_width) > max_size: + scale = max_size / max(final_height, final_width) + new_width = int(final_width * scale) + new_h1 = int(h1_w * scale) + new_h2 = int(h2_w * scale) + img1r = torch.nn.functional.interpolate(img1r.permute(0, 3, 1, 2), size=(new_h1, new_width), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + img2r = torch.nn.functional.interpolate(img2r.permute(0, 3, 1, 2), size=(new_h2, new_width), mode='bilinear', align_corners=False).permute(0, 2, 3, 1) + final_height = new_h1 + new_h2 + (gap if gap > 0 else 0) + final_width = new_width + if background_color == "transparent": + output = torch.zeros((1, final_height, final_width, img1.shape[-1]), dtype=img1.dtype, device=img1.device) + else: + color_value = 1.0 if background_color == "white" else 0.0 if background_color == "black" else 0.5 + output = torch.full((1, final_height, final_width, img1.shape[-1]), color_value, dtype=img1.dtype, device=img1.device) + output[:, :img1r.shape[1], :] = img1r + output[:, img1r.shape[1]+(gap if gap > 0 else 0):, :] = img2r return (output,) -# Node mappings NODE_CLASS_MAPPINGS = { "ImageConcatenateMulti_UTK": ImageConcatenateMulti_UTK, } diff --git a/nodes/image/imitation_hue_node.py b/nodes/image/imitation_hue_node.py index 407d20e..2016ce0 100644 --- a/nodes/image/imitation_hue_node.py +++ b/nodes/image/imitation_hue_node.py @@ -11,6 +11,12 @@ Performs color transfer and imitation between images with skin protection. import numpy as np import cv2 import torch +try: + from comfy.utils import ProgressBar +except ImportError: + ProgressBar = None +from ..image_utils import pil2tensor +from PIL import Image def image_stats(image): return np.mean(image[:, :, 1:], axis=(0, 1)), np.std(image[:, :, 1:], axis=(0, 1)) @@ -237,25 +243,29 @@ Performs color transfer and imitation between images with skin protection. def imitation_hue(self, imitation_image, target_image, strength, skin_protection, auto_brightness, brightness_range, auto_contrast, contrast_range, auto_saturation, saturation_range, auto_tone, tone_strength, mask=None): - for img in imitation_image: - img_cv1 = tensor2cv2(img) - - for img in target_image: + # 只取一张imitation_image + img_cv1 = tensor2cv2(imitation_image[0]) + results = [] + num_targets = len(target_image) + has_mask = mask is not None and len(mask) == num_targets + pb = ProgressBar(num_targets) if ProgressBar else None + for idx, img in enumerate(target_image): + if pb: + pb.update(idx+1) img_cv2 = tensor2cv2(img) - - img_cv3 = None - if mask is not None: - for img3 in mask: - img_cv3 = img3.cpu().numpy() + img_cv3 = None + if has_mask: + m = mask[idx] + img_cv3 = m.cpu().numpy() img_cv3 = (img_cv3 * 255).astype(np.uint8) - - result_img = color_transfer(img_cv1, img_cv2, img_cv3, strength, skin_protection, auto_brightness, - brightness_range,auto_contrast, contrast_range, auto_saturation, - saturation_range, auto_tone, tone_strength) - result_img = cv2.cvtColor(result_img, cv2.COLOR_BGR2RGB) - rst = torch.from_numpy(result_img.astype(np.float32) / 255.0).unsqueeze(0) - - return (rst,) + result_img = color_transfer(img_cv1, img_cv2, img_cv3, strength, skin_protection, auto_brightness, + brightness_range, auto_contrast, contrast_range, auto_saturation, + saturation_range, auto_tone, tone_strength) + result_img = cv2.cvtColor(result_img, cv2.COLOR_BGR2RGB) + pil_img = Image.fromarray(result_img) + rst = pil2tensor(pil_img) + results.append(rst) + return (torch.cat(results, dim=0),) # Node mappings NODE_CLASS_MAPPINGS = {