381 lines
14 KiB
Python
381 lines
14 KiB
Python
import torch
|
|
import numpy as np
|
|
from PIL import Image
|
|
import json
|
|
import os
|
|
|
|
class SmartBatchProcessor:
|
|
"""
|
|
A ComfyUI node for batch processing multiple image operations with templates.
|
|
"""
|
|
|
|
def __init__(self):
|
|
pass
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"operation_mode": (["stitch_horizontal", "stitch_vertical", "create_grid", "add_borders", "split_images"], {"default": "stitch_horizontal"}),
|
|
"batch_size": ("INT", {"default": 4, "min": 1, "max": 9}),
|
|
"output_format": (["individual", "combined"], {"default": "individual"}),
|
|
"save_template": ("BOOLEAN", {"default": False}),
|
|
"template_name": ("STRING", {"default": "my_template"}),
|
|
},
|
|
"optional": {
|
|
# Stitching parameters
|
|
"stitch_direction": (["left", "right", "top", "bottom"], {"default": "right"}),
|
|
"stitch_alignment": (["start", "center", "end"], {"default": "center"}),
|
|
"resize_mode": (["none", "largest", "smallest"], {"default": "none"}),
|
|
"spacing": ("INT", {"default": 10, "min": 0, "max": 100}),
|
|
|
|
# Grid parameters
|
|
"grid_rows": ("INT", {"default": 2, "min": 1, "max": 4}),
|
|
"grid_cols": ("INT", {"default": 2, "min": 1, "max": 4}),
|
|
|
|
# Border parameters
|
|
"border_width": ("INT", {"default": 20, "min": 0, "max": 100}),
|
|
"border_color": (["white", "black", "gray"], {"default": "white"}),
|
|
|
|
# Input images
|
|
"image_1": ("IMAGE",),
|
|
"image_2": ("IMAGE",),
|
|
"image_3": ("IMAGE",),
|
|
"image_4": ("IMAGE",),
|
|
"image_5": ("IMAGE",),
|
|
"image_6": ("IMAGE",),
|
|
"image_7": ("IMAGE",),
|
|
"image_8": ("IMAGE",),
|
|
"image_9": ("IMAGE",),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("IMAGE", "IMAGE", "IMAGE", "IMAGE", "STRING")
|
|
RETURN_NAMES = ("result_1", "result_2", "result_3", "result_4", "processing_log")
|
|
FUNCTION = "process_batch"
|
|
CATEGORY = "Niutonian/Image Processing"
|
|
|
|
def tensor_to_pil(self, tensor):
|
|
"""Convert tensor to PIL Image"""
|
|
if tensor is None:
|
|
return None
|
|
|
|
if len(tensor.shape) == 4:
|
|
tensor = tensor[0]
|
|
|
|
np_image = tensor.cpu().numpy()
|
|
|
|
if np_image.dtype == np.float32 or np_image.dtype == np.float64:
|
|
np_image = (np_image * 255).astype(np.uint8)
|
|
|
|
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)
|
|
return tensor
|
|
|
|
def get_valid_images(self, **kwargs):
|
|
"""Extract valid images from inputs"""
|
|
valid_images = []
|
|
|
|
for i in range(1, 10): # image_1 through image_9
|
|
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 create_empty_tensor(self):
|
|
"""Create empty tensor for unused outputs"""
|
|
empty_img = Image.new('RGB', (64, 64), (0, 0, 0))
|
|
return self.pil_to_tensor(empty_img)
|
|
|
|
def get_border_color(self, color_name):
|
|
"""Get border color from name"""
|
|
colors = {
|
|
"white": (255, 255, 255),
|
|
"black": (0, 0, 0),
|
|
"gray": (128, 128, 128),
|
|
}
|
|
return colors.get(color_name, (255, 255, 255))
|
|
|
|
def resize_image_to_match(self, images, resize_mode):
|
|
"""Resize images based on mode"""
|
|
if not images or resize_mode == "none":
|
|
return images
|
|
|
|
sizes = [(img.width, img.height) for img in images]
|
|
|
|
if resize_mode == "largest":
|
|
target_width = max(size[0] for size in sizes)
|
|
target_height = max(size[1] for size in sizes)
|
|
elif resize_mode == "smallest":
|
|
target_width = min(size[0] for size in sizes)
|
|
target_height = min(size[1] for size in sizes)
|
|
else:
|
|
return images
|
|
|
|
resized_images = []
|
|
for img in images:
|
|
if img.size != (target_width, target_height):
|
|
resized_img = img.resize((target_width, target_height), Image.LANCZOS)
|
|
resized_images.append(resized_img)
|
|
else:
|
|
resized_images.append(img)
|
|
|
|
return resized_images
|
|
|
|
def stitch_images_horizontal(self, images, alignment, spacing):
|
|
"""Stitch images horizontally"""
|
|
if not images:
|
|
return None
|
|
|
|
if len(images) == 1:
|
|
return images[0]
|
|
|
|
# Calculate canvas size
|
|
total_width = sum(img.width for img in images) + spacing * (len(images) - 1)
|
|
max_height = max(img.height for img in images)
|
|
|
|
# Create canvas
|
|
canvas = Image.new('RGB', (total_width, max_height), (255, 255, 255))
|
|
|
|
# Place images
|
|
current_x = 0
|
|
for img in images:
|
|
if alignment == "start":
|
|
y = 0
|
|
elif alignment == "center":
|
|
y = (max_height - img.height) // 2
|
|
else: # end
|
|
y = max_height - img.height
|
|
|
|
canvas.paste(img, (current_x, y))
|
|
current_x += img.width + spacing
|
|
|
|
return canvas
|
|
|
|
def stitch_images_vertical(self, images, alignment, spacing):
|
|
"""Stitch images vertically"""
|
|
if not images:
|
|
return None
|
|
|
|
if len(images) == 1:
|
|
return images[0]
|
|
|
|
# Calculate canvas size
|
|
max_width = max(img.width for img in images)
|
|
total_height = sum(img.height for img in images) + spacing * (len(images) - 1)
|
|
|
|
# Create canvas
|
|
canvas = Image.new('RGB', (max_width, total_height), (255, 255, 255))
|
|
|
|
# Place images
|
|
current_y = 0
|
|
for img in images:
|
|
if alignment == "start":
|
|
x = 0
|
|
elif alignment == "center":
|
|
x = (max_width - img.width) // 2
|
|
else: # end
|
|
x = max_width - img.width
|
|
|
|
canvas.paste(img, (x, current_y))
|
|
current_y += img.height + spacing
|
|
|
|
return canvas
|
|
|
|
def create_grid_from_images(self, images, rows, cols, spacing):
|
|
"""Create grid from images"""
|
|
if not images:
|
|
return None
|
|
|
|
# Pad images list to fill grid
|
|
while len(images) < rows * cols:
|
|
images.append(images[-1] if images else Image.new('RGB', (100, 100), (128, 128, 128)))
|
|
|
|
# Calculate cell size (use first image as reference)
|
|
cell_width = images[0].width
|
|
cell_height = images[0].height
|
|
|
|
# Calculate canvas size
|
|
canvas_width = cols * cell_width + (cols - 1) * spacing
|
|
canvas_height = rows * cell_height + (rows - 1) * spacing
|
|
|
|
# Create canvas
|
|
canvas = Image.new('RGB', (canvas_width, canvas_height), (255, 255, 255))
|
|
|
|
# Place images in grid
|
|
img_index = 0
|
|
for row in range(rows):
|
|
for col in range(cols):
|
|
if img_index >= len(images):
|
|
break
|
|
|
|
x = col * (cell_width + spacing)
|
|
y = row * (cell_height + spacing)
|
|
|
|
# Resize image to cell size if needed
|
|
img = images[img_index]
|
|
if img.size != (cell_width, cell_height):
|
|
img = img.resize((cell_width, cell_height), Image.LANCZOS)
|
|
|
|
canvas.paste(img, (x, y))
|
|
img_index += 1
|
|
|
|
return canvas
|
|
|
|
def add_border_to_image(self, image, border_width, border_color):
|
|
"""Add border to single image"""
|
|
if border_width == 0:
|
|
return image
|
|
|
|
new_width = image.width + 2 * border_width
|
|
new_height = image.height + 2 * border_width
|
|
|
|
bordered = Image.new('RGB', (new_width, new_height), border_color)
|
|
bordered.paste(image, (border_width, border_width))
|
|
|
|
return bordered
|
|
|
|
def split_image_into_parts(self, image, rows, cols):
|
|
"""Split image into grid parts"""
|
|
cell_width = image.width // cols
|
|
cell_height = image.height // rows
|
|
|
|
parts = []
|
|
for row in range(rows):
|
|
for col in range(cols):
|
|
x = col * cell_width
|
|
y = row * cell_height
|
|
|
|
part = image.crop((x, y, x + cell_width, y + cell_height))
|
|
parts.append(part)
|
|
|
|
return parts
|
|
|
|
def save_template_config(self, template_name, config):
|
|
"""Save processing template configuration"""
|
|
try:
|
|
# Create templates directory if it doesn't exist
|
|
templates_dir = os.path.join(os.path.dirname(__file__), "templates")
|
|
os.makedirs(templates_dir, exist_ok=True)
|
|
|
|
# Save template
|
|
template_path = os.path.join(templates_dir, f"{template_name}.json")
|
|
with open(template_path, 'w') as f:
|
|
json.dump(config, f, indent=2)
|
|
|
|
return f"Template saved: {template_path}"
|
|
except Exception as e:
|
|
return f"Failed to save template: {str(e)}"
|
|
|
|
def process_batch(self, operation_mode, batch_size, output_format, save_template, template_name,
|
|
stitch_direction="right", stitch_alignment="center", resize_mode="none", spacing=10,
|
|
grid_rows=2, grid_cols=2, border_width=20, border_color="white", **kwargs):
|
|
"""Main batch processing function"""
|
|
|
|
# Get valid images
|
|
valid_images = self.get_valid_images(**kwargs)
|
|
|
|
if not valid_images:
|
|
empty = self.create_empty_tensor()
|
|
return (empty, empty, empty, empty, "No valid images provided")
|
|
|
|
processing_log = []
|
|
results = []
|
|
|
|
try:
|
|
if operation_mode == "stitch_horizontal":
|
|
# Resize if needed
|
|
if resize_mode != "none":
|
|
valid_images = self.resize_image_to_match(valid_images, resize_mode)
|
|
|
|
# Stitch horizontally
|
|
result = self.stitch_images_horizontal(valid_images, stitch_alignment, spacing)
|
|
if result:
|
|
results.append(result)
|
|
processing_log.append(f"Stitched {len(valid_images)} images horizontally")
|
|
|
|
elif operation_mode == "stitch_vertical":
|
|
# Resize if needed
|
|
if resize_mode != "none":
|
|
valid_images = self.resize_image_to_match(valid_images, resize_mode)
|
|
|
|
# Stitch vertically
|
|
result = self.stitch_images_vertical(valid_images, stitch_alignment, spacing)
|
|
if result:
|
|
results.append(result)
|
|
processing_log.append(f"Stitched {len(valid_images)} images vertically")
|
|
|
|
elif operation_mode == "create_grid":
|
|
# Create grid
|
|
result = self.create_grid_from_images(valid_images, grid_rows, grid_cols, spacing)
|
|
if result:
|
|
results.append(result)
|
|
processing_log.append(f"Created {grid_rows}x{grid_cols} grid from {len(valid_images)} images")
|
|
|
|
elif operation_mode == "add_borders":
|
|
# Add borders to each image
|
|
border_color_rgb = self.get_border_color(border_color)
|
|
for i, img in enumerate(valid_images):
|
|
bordered = self.add_border_to_image(img, border_width, border_color_rgb)
|
|
results.append(bordered)
|
|
processing_log.append(f"Added border to image {i+1}")
|
|
|
|
elif operation_mode == "split_images":
|
|
# Split first image into parts
|
|
if valid_images:
|
|
parts = self.split_image_into_parts(valid_images[0], grid_rows, grid_cols)
|
|
results.extend(parts[:4]) # Limit to 4 results
|
|
processing_log.append(f"Split image into {len(parts)} parts")
|
|
|
|
# Save template if requested
|
|
if save_template:
|
|
config = {
|
|
"operation_mode": operation_mode,
|
|
"stitch_direction": stitch_direction,
|
|
"stitch_alignment": stitch_alignment,
|
|
"resize_mode": resize_mode,
|
|
"spacing": spacing,
|
|
"grid_rows": grid_rows,
|
|
"grid_cols": grid_cols,
|
|
"border_width": border_width,
|
|
"border_color": border_color,
|
|
}
|
|
save_result = self.save_template_config(template_name, config)
|
|
processing_log.append(save_result)
|
|
|
|
except Exception as e:
|
|
processing_log.append(f"Error during processing: {str(e)}")
|
|
|
|
# Convert results to tensors
|
|
result_tensors = []
|
|
for i in range(4):
|
|
if i < len(results):
|
|
result_tensors.append(self.pil_to_tensor(results[i]))
|
|
else:
|
|
result_tensors.append(self.create_empty_tensor())
|
|
|
|
log_text = "\n".join(processing_log)
|
|
|
|
return (*result_tensors, log_text)
|
|
|
|
# Node registration
|
|
NODE_CLASS_MAPPINGS = {
|
|
"SmartBatchProcessor": SmartBatchProcessor
|
|
}
|
|
|
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
"SmartBatchProcessor": "Smart Batch Processor"
|
|
} |