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