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:
@@ -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
|
||||
|
||||
# 批量放置图片
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user