Files

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"
}