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.
This commit is contained in:
Cyber Dick Lang
2025-07-02 20:47:56 +08:00
parent 33e49af04f
commit 12349affd6
3 changed files with 292 additions and 184 deletions
+19 -14
View File
@@ -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
# 批量放置图片
+246 -153
View File
@@ -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,
}
+27 -17
View File
@@ -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 = {