338 lines
14 KiB
Python
338 lines
14 KiB
Python
import torch
|
|
import numpy as np
|
|
from PIL import Image, ImageOps
|
|
import comfy.utils
|
|
|
|
class SmartImageStitch:
|
|
"""
|
|
A ComfyUI node that intelligently stitches multiple images together.
|
|
Automatically ignores empty/bypassed inputs and aligns remaining images
|
|
in the specified direction.
|
|
"""
|
|
|
|
def __init__(self):
|
|
pass
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"direction": (["left", "right", "top", "bottom"], {"default": "right"}),
|
|
"alignment": (["start", "center", "end"], {"default": "center"}),
|
|
"resize_mode": (["none", "largest", "smallest", "longest_side", "shortest_side"], {"default": "none"}),
|
|
"resize_method": (["lanczos", "bilinear", "bicubic", "nearest"], {"default": "lanczos"}),
|
|
"aspect_ratio_handling": (["pad", "crop", "stretch"], {"default": "pad"}),
|
|
"spacing": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}),
|
|
"background_color": (["transparent", "white", "black"], {"default": "transparent"}),
|
|
},
|
|
"optional": {
|
|
"image_1": ("IMAGE",),
|
|
"image_2": ("IMAGE",),
|
|
"image_3": ("IMAGE",),
|
|
"image_4": ("IMAGE",),
|
|
"image_5": ("IMAGE",),
|
|
"image_6": ("IMAGE",),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("IMAGE",)
|
|
RETURN_NAMES = ("stitched_image",)
|
|
FUNCTION = "stitch_images"
|
|
CATEGORY = "Niutonian/Image Processing"
|
|
|
|
def tensor_to_pil(self, tensor):
|
|
"""Convert tensor to PIL Image"""
|
|
if tensor is None:
|
|
return None
|
|
|
|
# Handle batch dimension
|
|
if len(tensor.shape) == 4:
|
|
tensor = tensor[0]
|
|
|
|
# Convert from torch tensor to numpy
|
|
np_image = tensor.cpu().numpy()
|
|
|
|
# Convert from float [0,1] to uint8 [0,255]
|
|
if np_image.dtype == np.float32 or np_image.dtype == np.float64:
|
|
np_image = (np_image * 255).astype(np.uint8)
|
|
|
|
# Convert to PIL Image
|
|
if len(np_image.shape) == 3:
|
|
return Image.fromarray(np_image, 'RGB')
|
|
else:
|
|
return Image.fromarray(np_image, 'L')
|
|
|
|
def pil_to_tensor(self, pil_image):
|
|
"""Convert PIL Image to tensor"""
|
|
if pil_image.mode != 'RGB':
|
|
pil_image = pil_image.convert('RGB')
|
|
|
|
np_image = np.array(pil_image).astype(np.float32) / 255.0
|
|
tensor = torch.from_numpy(np_image).unsqueeze(0) # Add batch dimension
|
|
return tensor
|
|
|
|
def get_resize_method(self, method_name):
|
|
"""Get PIL resize method from string"""
|
|
methods = {
|
|
"lanczos": Image.LANCZOS,
|
|
"bilinear": Image.BILINEAR,
|
|
"bicubic": Image.BICUBIC,
|
|
"nearest": Image.NEAREST
|
|
}
|
|
return methods.get(method_name, Image.LANCZOS)
|
|
|
|
def calculate_target_size(self, images, resize_mode, direction):
|
|
"""Calculate target size for all images based on resize mode"""
|
|
if not images or resize_mode == "none":
|
|
return None
|
|
|
|
sizes = [(img.width, img.height) for img in images]
|
|
|
|
if resize_mode == "largest":
|
|
# Resize all to match the largest image dimensions
|
|
max_width = max(size[0] for size in sizes)
|
|
max_height = max(size[1] for size in sizes)
|
|
return (max_width, max_height)
|
|
|
|
elif resize_mode == "smallest":
|
|
# Resize all to match the smallest image dimensions
|
|
min_width = min(size[0] for size in sizes)
|
|
min_height = min(size[1] for size in sizes)
|
|
return (min_width, min_height)
|
|
|
|
elif resize_mode == "longest_side":
|
|
# Find the longest side across all images and make all images square to that size
|
|
max_dimension = max(max(size) for size in sizes)
|
|
return (max_dimension, max_dimension)
|
|
|
|
elif resize_mode == "shortest_side":
|
|
# Find the shortest side across all images and make all images square to that size
|
|
min_dimension = min(min(size) for size in sizes)
|
|
return (min_dimension, min_dimension)
|
|
|
|
return None
|
|
|
|
def resize_images_to_target(self, images, target_size, resize_method, aspect_handling):
|
|
"""Resize all images to target size while handling aspect ratios"""
|
|
if target_size is None:
|
|
return images
|
|
|
|
pil_method = self.get_resize_method(resize_method)
|
|
resized_images = []
|
|
target_width, target_height = target_size
|
|
|
|
for img in images:
|
|
if img.size == target_size:
|
|
resized_images.append(img)
|
|
continue
|
|
|
|
original_width, original_height = img.size
|
|
|
|
if aspect_handling == "stretch":
|
|
# Simple stretch (original behavior)
|
|
resized_img = img.resize(target_size, pil_method)
|
|
resized_images.append(resized_img)
|
|
|
|
elif aspect_handling == "pad":
|
|
# Maintain aspect ratio and pad with background
|
|
# Calculate scale to fit within target size
|
|
scale_w = target_width / original_width
|
|
scale_h = target_height / original_height
|
|
scale = min(scale_w, scale_h) # Use smaller scale to fit within bounds
|
|
|
|
# Calculate new size maintaining aspect ratio
|
|
new_width = int(original_width * scale)
|
|
new_height = int(original_height * scale)
|
|
|
|
# Resize image maintaining aspect ratio
|
|
resized_img = img.resize((new_width, new_height), pil_method)
|
|
|
|
# Create canvas with target size and paste resized image in center
|
|
if img.mode == 'RGBA':
|
|
canvas = Image.new('RGBA', target_size, (255, 255, 255, 0))
|
|
else:
|
|
canvas = Image.new('RGB', target_size, (255, 255, 255))
|
|
|
|
# Calculate position to center the image
|
|
paste_x = (target_width - new_width) // 2
|
|
paste_y = (target_height - new_height) // 2
|
|
|
|
if resized_img.mode == 'RGBA':
|
|
canvas.paste(resized_img, (paste_x, paste_y), resized_img)
|
|
else:
|
|
canvas.paste(resized_img, (paste_x, paste_y))
|
|
|
|
resized_images.append(canvas)
|
|
|
|
elif aspect_handling == "crop":
|
|
# Maintain aspect ratio and crop to fill target size
|
|
# Calculate scale to fill target size
|
|
scale_w = target_width / original_width
|
|
scale_h = target_height / original_height
|
|
scale = max(scale_w, scale_h) # Use larger scale to fill bounds
|
|
|
|
# Calculate new size maintaining aspect ratio
|
|
new_width = int(original_width * scale)
|
|
new_height = int(original_height * scale)
|
|
|
|
# Resize image maintaining aspect ratio
|
|
resized_img = img.resize((new_width, new_height), pil_method)
|
|
|
|
# Calculate crop box to center the crop
|
|
crop_x = (new_width - target_width) // 2
|
|
crop_y = (new_height - target_height) // 2
|
|
|
|
# Crop to target size
|
|
cropped_img = resized_img.crop((
|
|
crop_x,
|
|
crop_y,
|
|
crop_x + target_width,
|
|
crop_y + target_height
|
|
))
|
|
|
|
resized_images.append(cropped_img)
|
|
|
|
return resized_images
|
|
|
|
def get_valid_images(self, **kwargs):
|
|
"""Extract valid (non-None) images from inputs"""
|
|
valid_images = []
|
|
|
|
# Check each possible image input
|
|
for i in range(1, 7): # image_1 through image_6
|
|
image_key = f"image_{i}"
|
|
if image_key in kwargs and kwargs[image_key] is not None:
|
|
pil_img = self.tensor_to_pil(kwargs[image_key])
|
|
if pil_img is not None:
|
|
valid_images.append(pil_img)
|
|
|
|
return valid_images
|
|
|
|
def calculate_canvas_size(self, images, direction, spacing):
|
|
"""Calculate the size needed for the final canvas"""
|
|
if not images:
|
|
return (100, 100) # Default size if no images
|
|
|
|
if direction in ["left", "right"]:
|
|
# Horizontal stitching
|
|
total_width = sum(img.width for img in images) + spacing * (len(images) - 1)
|
|
max_height = max(img.height for img in images)
|
|
return (total_width, max_height)
|
|
else:
|
|
# Vertical stitching
|
|
max_width = max(img.width for img in images)
|
|
total_height = sum(img.height for img in images) + spacing * (len(images) - 1)
|
|
return (max_width, total_height)
|
|
|
|
def get_background_color(self, bg_type, mode='RGB'):
|
|
"""Get background color based on type"""
|
|
if bg_type == "white":
|
|
return (255, 255, 255) if mode == 'RGB' else 255
|
|
elif bg_type == "black":
|
|
return (0, 0, 0) if mode == 'RGB' else 0
|
|
else: # transparent
|
|
return (255, 255, 255, 0) if mode == 'RGBA' else (255, 255, 255)
|
|
|
|
def calculate_position(self, img_size, canvas_size, alignment, direction, current_offset):
|
|
"""Calculate position for placing an image on the canvas"""
|
|
img_width, img_height = img_size
|
|
canvas_width, canvas_height = canvas_size
|
|
|
|
if direction in ["left", "right"]:
|
|
# Horizontal stitching
|
|
x = current_offset
|
|
if alignment == "start":
|
|
y = 0
|
|
elif alignment == "center":
|
|
y = (canvas_height - img_height) // 2
|
|
else: # end
|
|
y = canvas_height - img_height
|
|
else:
|
|
# Vertical stitching
|
|
y = current_offset
|
|
if alignment == "start":
|
|
x = 0
|
|
elif alignment == "center":
|
|
x = (canvas_width - img_width) // 2
|
|
else: # end
|
|
x = canvas_width - img_width
|
|
|
|
return (x, y)
|
|
|
|
def stitch_images(self, direction, alignment, resize_mode, resize_method, aspect_ratio_handling, spacing, background_color, **kwargs):
|
|
"""Main function to stitch images together"""
|
|
|
|
# Get all valid images
|
|
valid_images = self.get_valid_images(**kwargs)
|
|
|
|
if not valid_images:
|
|
# Return a small default image if no valid images
|
|
default_img = Image.new('RGB', (100, 100), (128, 128, 128))
|
|
return (self.pil_to_tensor(default_img),)
|
|
|
|
if len(valid_images) == 1:
|
|
# Return single image if only one valid image
|
|
return (self.pil_to_tensor(valid_images[0]),)
|
|
|
|
# Apply resizing if specified
|
|
target_size = self.calculate_target_size(valid_images, resize_mode, direction)
|
|
if target_size:
|
|
valid_images = self.resize_images_to_target(valid_images, target_size, resize_method, aspect_ratio_handling)
|
|
|
|
# Reverse order for left and top directions to maintain logical flow
|
|
if direction in ["left", "top"]:
|
|
valid_images = valid_images[::-1]
|
|
|
|
# Calculate canvas size
|
|
canvas_size = self.calculate_canvas_size(valid_images, direction, spacing)
|
|
|
|
# Create canvas
|
|
if background_color == "transparent":
|
|
canvas = Image.new('RGBA', canvas_size, self.get_background_color(background_color, 'RGBA'))
|
|
else:
|
|
canvas = Image.new('RGB', canvas_size, self.get_background_color(background_color, 'RGB'))
|
|
|
|
# Place images on canvas
|
|
current_offset = 0
|
|
|
|
for i, img in enumerate(valid_images):
|
|
# Convert to same mode as canvas if needed
|
|
if canvas.mode == 'RGBA' and img.mode != 'RGBA':
|
|
img = img.convert('RGBA')
|
|
elif canvas.mode == 'RGB' and img.mode == 'RGBA':
|
|
# Create a white background and paste the RGBA image onto it
|
|
white_bg = Image.new('RGB', img.size, (255, 255, 255))
|
|
white_bg.paste(img, mask=img.split()[-1] if img.mode == 'RGBA' else None)
|
|
img = white_bg
|
|
|
|
# Calculate position
|
|
pos = self.calculate_position(img.size, canvas_size, alignment, direction, current_offset)
|
|
|
|
# Paste image
|
|
if img.mode == 'RGBA':
|
|
canvas.paste(img, pos, img)
|
|
else:
|
|
canvas.paste(img, pos)
|
|
|
|
# Update offset for next image
|
|
if direction in ["left", "right"]:
|
|
current_offset += img.width + spacing
|
|
else:
|
|
current_offset += img.height + spacing
|
|
|
|
# Convert back to RGB if it was RGBA but background wasn't transparent
|
|
if canvas.mode == 'RGBA' and background_color != "transparent":
|
|
rgb_canvas = Image.new('RGB', canvas.size, self.get_background_color(background_color, 'RGB'))
|
|
rgb_canvas.paste(canvas, mask=canvas.split()[-1])
|
|
canvas = rgb_canvas
|
|
|
|
return (self.pil_to_tensor(canvas),)
|
|
|
|
# Node registration
|
|
NODE_CLASS_MAPPINGS = {
|
|
"SmartImageStitch": SmartImageStitch
|
|
}
|
|
|
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
"SmartImageStitch": "Smart Image Stitch"
|
|
} |