scale down fix

This commit is contained in:
dengjia
2024-11-14 15:34:26 +08:00
parent c6e0126fcf
commit ff606a511c
2 changed files with 198 additions and 45 deletions
+40 -24
View File
@@ -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,)
+158 -21
View File
@@ -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):