scale down fix
This commit is contained in:
@@ -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,)
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user