diff --git a/nodes.py b/nodes.py index 31ed439..3ac09af 100644 --- a/nodes.py +++ b/nodes.py @@ -2,6 +2,7 @@ import torch import numpy as np from .utils import resize_pixel_art, convert_to_grayscale, convert_to_bw from .utils import PaletteGenerator, Dithering +import time class ComfyUIPixelArtAdvanced: """ @@ -16,9 +17,10 @@ class ComfyUIPixelArtAdvanced: "downscale_factor": ("INT", { "default": 4, "min": 1, - "max": 16, + "max": 32, "step": 1 }), + "scale_mode": (["auto", "nearest", "area", "linear", "cubic", "lanczos"],), "rescale_to_original": ("BOOLEAN", {"default": False}), "color_mode": (["rgb", "grayscale", "bw"],), "colors": ("INT", { @@ -27,15 +29,13 @@ class ComfyUIPixelArtAdvanced: "max": 256, "step": 1 }), - "quantization_method": (["kmeans", "median_cut"],), + "quantization_method": (["auto", "kmeans", "mediancut", "maxcoverage", "fastoctree", "libimagequant", "median_cut"],), "dithering": (["none", "floyd-steinberg"],), - "palette_type": (["adaptive", "custom", "from_image"],), }, "optional": { - "custom_palette": ("STRING", {"default": "15,56,15;48,98,48;139,172,15;155,188,15"}), "palette_image": ("IMAGE",), "palette_size": ("INT", { - "default": 16, + "default": 32, "min": 2, "max": 256, "step": 1 @@ -47,49 +47,65 @@ class ComfyUIPixelArtAdvanced: FUNCTION = "process" CATEGORY = "image/Pixel Art" - def process(self, image, downscale_factor, rescale_to_original, color_mode, colors, - quantization_method, dithering, palette_type, - custom_palette=None, palette_image=None, palette_size=16): + def process(self, image, downscale_factor, scale_mode, + rescale_to_original, color_mode, colors, + quantization_method, dithering, palette_image=None, palette_size=16): + start_time = time.time() + # Convert from torch tensor to numpy image_np = image[0].cpu().numpy() image_np = (image_np * 255).astype(np.uint8) + print(f"Conversion to numpy: {time.time() - start_time:.2f}s ") + # Store original size if rescaling is needed original_size = (image_np.shape[1], image_np.shape[0]) if rescale_to_original else None - # Apply pixel art scaling - image_np = resize_pixel_art(image_np, downscale_factor, rescale_to_original, original_size) + # Apply pixel art scaling with new parameters + t0 = time.time() + image_np = resize_pixel_art( + image_np, + downscale_factor, + rescale_to_original=rescale_to_original, + original_size=original_size, + scale_down_mode=scale_mode + ) + print(f"Pixel art scaling ({scale_mode}): {time.time() - t0:.2f}s ") # Apply color mode conversion + t0 = time.time() if color_mode == "grayscale": image_np = convert_to_grayscale(image_np) elif color_mode == "bw": image_np = convert_to_bw(image_np) + print(f"Color mode conversion: {time.time() - t0:.2f}s ") # Get palette - if palette_type == "custom" and custom_palette: - palette = PaletteGenerator.parse_custom_palette(custom_palette) - if palette is None: - palette = PaletteGenerator.kmeans_palette(image_np, colors) - elif palette_type == "from_image" and palette_image is not None: + t0 = time.time() + if palette_image is not None: # Convert palette image from torch tensor to numpy palette_np = palette_image[0].cpu().numpy() palette_np = (palette_np * 255).astype(np.uint8) - palette = PaletteGenerator.kmeans_palette(palette_np, palette_size) + palette = PaletteGenerator.get_palette(palette_np, palette_size) else: # adaptive - if quantization_method == "kmeans": - palette = PaletteGenerator.kmeans_palette(image_np, colors) - else: # median_cut - palette = PaletteGenerator.median_cut_palette(image_np, colors) + print(f"get_palette use {quantization_method}") + palette = PaletteGenerator.get_palette(image_np, colors, quantization_method) + + print(f"Palette generation: {time.time() - t0:.2f}s ") - # Apply dithering and quantization + # Apply dithering if requested + t0 = time.time() if dithering == "floyd-steinberg": - result = Dithering.floyd_steinberg(image_np, palette) + image_np = Dithering.floyd_steinberg(image_np, palette) else: - result = Dithering.simple_quantize(image_np, palette) + image_np = Dithering.simple_quantize(image_np, palette) + print(f"Dithering: {time.time() - t0:.2f}s") # Convert back to torch tensor - result = torch.from_numpy(result.astype(np.float32) / 255.0) + result = torch.from_numpy(image_np.astype(np.float32) / 255.0) result = result.unsqueeze(0) + print(f"Conversion to tensor: {time.time() - t0:.2f}s ") + + print(f"Total processing time: {time.time() - start_time:.2f}s ") return (result,) \ No newline at end of file diff --git a/utils.py b/utils.py index 4259d50..7217d42 100644 --- a/utils.py +++ b/utils.py @@ -3,17 +3,38 @@ import cv2 from PIL import Image import colorsys -def scale_down(image, scale_factor): +def scale_down(image, scale_factor, mode='auto'): """ - Downscale image by integer factor using area interpolation + Downscale image with advanced interpolation options Args: image: numpy array (H, W, C) scale_factor: int, factor to reduce image by + mode: str, interpolation mode: + - 'auto': Automatically select best method + - 'area': cv2.INTER_AREA (good for downscaling) + - 'nearest': cv2.INTER_NEAREST (preserves exact colors) + - 'linear': cv2.INTER_LINEAR (smooth but can blur) + - 'cubic': cv2.INTER_CUBIC (sharper than linear) + - 'lanczos': cv2.INTER_LANCZOS4 (high quality, can preserve edges) """ h, w = image.shape[:2] small_h, small_w = h // scale_factor, w // scale_factor - return cv2.resize(image, (small_w, small_h), interpolation=cv2.INTER_AREA) + + # Dictionary of interpolation methods + interpolation_methods = { + 'nearest': cv2.INTER_NEAREST, # Best for pixel art + 'area': cv2.INTER_AREA, # Good general downscaling + 'linear': cv2.INTER_LINEAR, # Smooth, can blur + 'cubic': cv2.INTER_CUBIC, # Sharper edges + 'lanczos': cv2.INTER_LANCZOS4 # High quality + } + if mode == 'auto': + mode = 'area' + + # Apply selected interpolation + return cv2.resize(image, (small_w, small_h), + interpolation=interpolation_methods.get(mode, cv2.INTER_AREA)) def scale_up(image, scale_factor): """ @@ -27,22 +48,38 @@ def scale_up(image, scale_factor): large_h, large_w = h * scale_factor, w * scale_factor return cv2.resize(image, (large_w, large_h), interpolation=cv2.INTER_NEAREST) -def resize_pixel_art(image, scale_factor, rescale_to_original=False, original_size=None): +def resize_pixel_art(image, scale_factor, rescale_to_original=False, original_size=None, + scale_down_mode='auto'): """ - Resize image using pixel art scaling + Resize image using advanced pixel art scaling methods Args: image: numpy array (H, W, C) scale_factor: int, factor to reduce image by rescale_to_original: bool, whether to rescale back to original size original_size: tuple (width, height), original image dimensions + scale_down_mode: str, interpolation mode: + - 'auto': Automatically select best method + - 'nearest': Best for pixel art + - 'area': Good for general downscaling + - 'linear': Smooth but can blur + - 'cubic': Sharper edges + - 'lanczos': High quality + Returns: + numpy array: Resized image """ - # Scale down first - downscaled = scale_down(image, scale_factor) + # Scale down first using advanced methods + downscaled = scale_down( + image, + scale_factor, + scale_down_mode + ) # If rescaling is requested and we have original size if rescale_to_original and original_size: - return cv2.resize(downscaled, original_size, interpolation=cv2.INTER_NEAREST) + h, w = original_size + # Always use nearest neighbor for upscaling to maintain pixel art look + return cv2.resize(downscaled, (w, h), interpolation=cv2.INTER_NEAREST) return downscaled @@ -58,17 +95,127 @@ def convert_to_bw(image, threshold=127): return cv2.cvtColor(bw, cv2.COLOR_GRAY2RGB) class PaletteGenerator: + @staticmethod + def pillow_palette(image, n_colors, method='libimagequant'): + """Generate palette using Pillow's quantization methods + + Methods: + - 'libimagequant': Highest quality, reasonable speed (default) + - 'mediancut': Fast, good quality + - 'maxcoverage': Better color coverage + - 'fastoctree': Fastest, slightly lower quality + """ + # Convert numpy array to PIL Image + if isinstance(image, np.ndarray): + image = Image.fromarray(image) + + # Convert to RGB mode if needed + if image.mode != 'RGB': + image = image.convert('RGB') + + # Quantize image + if method == 'mediancut': + quantized = image.quantize(colors=n_colors, method=Image.Quantize.MEDIANCUT) + elif method == 'maxcoverage': + quantized = image.quantize(colors=n_colors, method=Image.Quantize.MAXCOVERAGE) + elif method == 'fastoctree': + quantized = image.quantize(colors=n_colors, method=Image.Quantize.FASTOCTREE) + else: # libimagequant (default) + quantized = image.quantize(colors=n_colors, method=Image.Quantize.LIBIMAGEQUANT) + + # Extract palette + palette = np.array(quantized.getpalette()[:n_colors*3]).reshape(-1, 3) + return palette + + @staticmethod + def get_palette(image, n_colors, method='auto'): + """Smart palette generation using multiple methods + + Methods: + - 'auto': Choose best method based on image size and n_colors + - 'libimagequant': Pillow's high quality quantizer + - 'mediancut': Pillow's median cut + - 'maxcoverage': Pillow's maximum coverage + - 'fastoctree': Pillow's fast octree + - 'kmeans': OpenCV k-means clustering + - 'median_cut': Custom median cut implementation + """ + # Auto method selection + if method == 'auto': + image_size = image.shape[0] * image.shape[1] + if image_size > 1000000 or n_colors > 32: # Large image or many colors + method = 'fastoctree' + elif image_size > 500000: # Medium image + method = 'libimagequant' + else: # Small image + method = 'kmeans' + + print(f'Using method: {method} image size: {image.shape[0] * image.shape[1]} ') + # Use Pillow methods first + if method in ['libimagequant', 'mediancut', 'maxcoverage', 'fastoctree']: + return PaletteGenerator.pillow_palette(image, n_colors, method) + + # Fallback to OpenCV k-means + elif method == 'kmeans': + return PaletteGenerator.kmeans_palette(image, n_colors) + + # Custom median cut implementation + elif method == 'median_cut': + return PaletteGenerator.median_cut_palette(image, n_colors) + + else: + raise ValueError(f"Unknown method: {method}") + @staticmethod def kmeans_palette(image, n_colors): - """Generate palette using k-means clustering""" + """Fallback k-means clustering method""" + # Limit maximum colors for performance + n_colors = min(n_colors, 32) pixels = image.reshape(-1, 3).astype(np.float32) criteria = (cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, 200, 0.1) - _, _, palette = cv2.kmeans(pixels, n_colors, None, criteria, 10, cv2.KMEANS_RANDOM_CENTERS) + + try: + # Check if CUDA is available + if cv2.cuda.getCudaEnabledDeviceCount() > 0: + # Move data to GPU + gpu_mat = cv2.cuda_GpuMat() + gpu_mat.upload(pixels) + + # Run k-means with GPU acceleration + _, _, palette = cv2.kmeans( + gpu_mat.download(), + n_colors, + None, + criteria, + 10, + cv2.KMEANS_RANDOM_CENTERS + cv2.KMEANS_USE_GPU + ) + print("Using GPU acceleration for k-means") + return palette + + except Exception as e: + print(f"GPU acceleration failed, falling back to CPU: {str(e)}") + + # CPU fallback with data reduction for better performance + if len(pixels) > 10000: + indices = np.random.choice(len(pixels), 10000, replace=False) + pixels = pixels[indices] + + # Run k-means on CPU + _, _, palette = cv2.kmeans( + pixels, + n_colors, + None, + criteria, + 10, + cv2.KMEANS_RANDOM_CENTERS + ) + print("Using CPU for k-means") return palette @staticmethod def median_cut_palette(image, n_colors): - """Generate palette using median cut algorithm""" + """Fallback median cut implementation""" pixels = image.reshape(-1, 3) def cut_box(box_pixels): @@ -87,16 +234,6 @@ class PaletteGenerator: boxes.extend([box1, box2]) return np.array([np.mean(box, axis=0) for box in boxes]) - - @staticmethod - def parse_custom_palette(palette_str): - """Parse custom palette string in format 'R,G,B;R,G,B;...'""" - try: - return np.array([list(map(int, color.split(','))) - for color in palette_str.split(';')]) - except: - return None - class Dithering: @staticmethod def extract_palette_from_image(image, palette_size=16):