Files

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