From 462dbc9bb368a384f601fc46180aa9d96a471824 Mon Sep 17 00:00:00 2001 From: Fillip Date: Wed, 14 Aug 2024 15:54:34 -0700 Subject: [PATCH] Upscale Model node added --- __init__.py | 5 +- nodes/{FL_PixelShader.py => FL_PixelArt.py} | 244 ++++++++++---------- nodes/FL_UpscaleModel.py | 74 ++++++ 3 files changed, 200 insertions(+), 123 deletions(-) rename nodes/{FL_PixelShader.py => FL_PixelArt.py} (95%) create mode 100644 nodes/FL_UpscaleModel.py diff --git a/__init__.py b/__init__.py index 15284b0..a0bec57 100644 --- a/__init__.py +++ b/__init__.py @@ -14,7 +14,7 @@ from .nodes.FL_HalfTone import FL_HalftonePattern from .nodes.FL_RandomRange import FL_RandomNumber from .nodes.FL_PromptSelector import FL_PromptSelector from .nodes.FL_Shader import FL_Shadertoy -from .nodes.FL_PixelShader import FL_PixelArtShader +from .nodes.FL_PixelArt import FL_PixelArtShader from .nodes.FL_InfiniteZoom import FL_InfiniteZoom from .nodes.FL_PaperDrawn import FL_PaperDrawn from .nodes.FL_ImageNotes import FL_ImageNotes @@ -50,6 +50,7 @@ from .nodes.FL_CaptionToCSV import FL_CaptionToCSV from .nodes.FL_KsamplerPlus import FL_KsamplerPlus from .nodes.FL_KsamplerBasic import FL_KsamplerBasic from .nodes.FL_KsamplerFractals import FL_FractalKSampler +from .nodes.FL_UpscaleModel import FL_UpscaleModel @@ -107,6 +108,7 @@ NODE_CLASS_MAPPINGS = { "FL_KsamplerPlus": FL_KsamplerPlus, "FL_KsamplerBasic": FL_KsamplerBasic, "FL_FractalKSampler": FL_FractalKSampler, + "FL_UpscaleModel": FL_UpscaleModel, } @@ -163,6 +165,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FL_KsamplerPlus": "FL KSampler Plus", "FL_KsamplerBasic": "FL KSampler Basic", "FL_FractalKSampler": "FL Fractal KSampler", + "FL_UpscaleModel": "FL Upscale Model", } diff --git a/nodes/FL_PixelShader.py b/nodes/FL_PixelArt.py similarity index 95% rename from nodes/FL_PixelShader.py rename to nodes/FL_PixelArt.py index 8e37e38..2948d6b 100644 --- a/nodes/FL_PixelShader.py +++ b/nodes/FL_PixelArt.py @@ -1,123 +1,123 @@ -import torch -import numpy as np -from PIL import Image -from sklearn.cluster import KMeans - -from comfy.utils import ProgressBar - -class FL_PixelArtShader: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "images": ("IMAGE",), - }, - "optional": { - "pixel_size": ("FLOAT", {"default": 100.0, "min": 1.0, "max": 1000.0, "step": 1.0}), - "color_depth": ("FLOAT", {"default": 50.0, "min": 1.0, "max": 255.0, "step": 1.0}), - "use_aspect_ratio": ("BOOLEAN", {"default": True}), - "palette_image": ("IMAGE", {"default": None}), - "palette_colors": ("INT", {"default": 16, "min": 2, "max": 15, "step": 1}), - "mask": ("IMAGE", {"default": None}), - }, - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "apply_pixel_art_shader" - CATEGORY = "🏵️Fill Nodes/VFX" - - def apply_pixel_art_shader(self, images, use_aspect_ratio, pixel_size, color_depth, palette_image=None, - palette_colors=16, mask=None): - result = [] - total_images = len(images) - pbar = ProgressBar(total_images) - - if palette_image is not None: - palette = extract_palette(self.t2p(palette_image[0]), palette_colors) - else: - palette = None - - mask_images = self.prepare_mask_batch(mask, total_images) if mask is not None else None - - for idx, image in enumerate(images): - img = self.t2p(image) - - mask_img = self.process_mask(mask_images[idx], img.size) if mask_images is not None else None - - result_img = pixel_art_effect(img, pixel_size, color_depth, use_aspect_ratio, palette, mask_img) - result_img = self.p2t(result_img) - result.append(result_img) - pbar.update_absolute(idx + 1) - - return (torch.cat(result, dim=0),) - - def t2p(self, t): - i = 255.0 * t.cpu().numpy().squeeze() - return Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) - - def p2t(self, p): - i = np.array(p).astype(np.float32) / 255.0 - return torch.from_numpy(i).unsqueeze(0) - - def prepare_mask_batch(self, mask, total_images): - if mask is None: - return None - mask_images = [self.t2p(m) for m in mask] - if len(mask_images) < total_images: - mask_images = mask_images * (total_images // len(mask_images) + 1) - return mask_images[:total_images] - - def process_mask(self, mask, target_size): - mask = mask.resize(target_size, Image.LANCZOS) - return mask.convert('L') if mask.mode != 'L' else mask - -def extract_palette(image, n_colors): - image = image.convert('RGB') - pixels = np.array(image).reshape(-1, 3) - kmeans = KMeans(n_clusters=n_colors, random_state=42) - kmeans.fit(pixels) - colors = kmeans.cluster_centers_ - return torch.from_numpy(colors.astype(np.float32) / 255.0).to("cuda") - -def pixel_art_effect(image, pixel_size, color_depth, use_aspect_ratio, palette, mask=None): - image = torch.tensor(np.array(image)).float().to("cuda") / 255.0 - height, width = image.shape[0], image.shape[1] - - if use_aspect_ratio: - aspect_ratio = width / height - pixel_size_x, pixel_size_y = pixel_size, pixel_size / aspect_ratio - else: - pixel_size_x = pixel_size_y = pixel_size - - new_width = int(width / pixel_size_x) - new_height = int(height / pixel_size_y) - - # Resize the image to create the pixelated effect - pixelated = image.permute(2, 0, 1).unsqueeze(0) - pixelated = torch.nn.functional.interpolate(pixelated, size=(new_height, new_width), mode='nearest') - pixelated = torch.nn.functional.interpolate(pixelated, size=(height, width), mode='nearest') - pixelated = pixelated.squeeze(0).permute(1, 2, 0) - - # Apply color depth reduction - pixelated = adjust_color(pixelated, color_depth) - - if palette is not None: - pixelated = apply_palette(pixelated, palette) - - if mask is not None: - mask_tensor = torch.tensor(np.array(mask)).float().to("cuda") / 255.0 - mask_tensor = mask_tensor.unsqueeze(-1).expand(-1, -1, 3) - pixelated = pixelated * mask_tensor + image * (1 - mask_tensor) - - return Image.fromarray((pixelated.cpu().numpy() * 255).astype(np.uint8)) - -def adjust_color(color, color_depth): - return torch.floor(color * color_depth) / color_depth - -def apply_palette(image, palette): - original_shape = image.shape - pixels = image.reshape(-1, 3) - distances = torch.cdist(pixels, palette) - nearest_palette_indices = torch.argmin(distances, dim=1) - new_pixels = palette[nearest_palette_indices] +import torch +import numpy as np +from PIL import Image +from sklearn.cluster import KMeans + +from comfy.utils import ProgressBar + +class FL_PixelArtShader: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "images": ("IMAGE",), + }, + "optional": { + "pixel_size": ("FLOAT", {"default": 15.0, "min": 1.0, "max": 100.0, "step": 1.0}), + "color_depth": ("FLOAT", {"default": 50.0, "min": 1.0, "max": 255.0, "step": 1.0}), + "use_aspect_ratio": ("BOOLEAN", {"default": True}), + "palette_image": ("IMAGE", {"default": None}), + "palette_colors": ("INT", {"default": 16, "min": 2, "max": 15, "step": 1}), + "mask": ("IMAGE", {"default": None}), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "apply_pixel_art_shader" + CATEGORY = "🏵️Fill Nodes/VFX" + + def apply_pixel_art_shader(self, images, use_aspect_ratio, pixel_size, color_depth, palette_image=None, + palette_colors=16, mask=None): + result = [] + total_images = len(images) + pbar = ProgressBar(total_images) + + if palette_image is not None: + palette = extract_palette(self.t2p(palette_image[0]), palette_colors) + else: + palette = None + + mask_images = self.prepare_mask_batch(mask, total_images) if mask is not None else None + + for idx, image in enumerate(images): + img = self.t2p(image) + + mask_img = self.process_mask(mask_images[idx], img.size) if mask_images is not None else None + + result_img = pixel_art_effect(img, pixel_size, color_depth, use_aspect_ratio, palette, mask_img) + result_img = self.p2t(result_img) + result.append(result_img) + pbar.update_absolute(idx + 1) + + return (torch.cat(result, dim=0),) + + def t2p(self, t): + i = 255.0 * t.cpu().numpy().squeeze() + return Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)) + + def p2t(self, p): + i = np.array(p).astype(np.float32) / 255.0 + return torch.from_numpy(i).unsqueeze(0) + + def prepare_mask_batch(self, mask, total_images): + if mask is None: + return None + mask_images = [self.t2p(m) for m in mask] + if len(mask_images) < total_images: + mask_images = mask_images * (total_images // len(mask_images) + 1) + return mask_images[:total_images] + + def process_mask(self, mask, target_size): + mask = mask.resize(target_size, Image.LANCZOS) + return mask.convert('L') if mask.mode != 'L' else mask + +def extract_palette(image, n_colors): + image = image.convert('RGB') + pixels = np.array(image).reshape(-1, 3) + kmeans = KMeans(n_clusters=n_colors, random_state=42) + kmeans.fit(pixels) + colors = kmeans.cluster_centers_ + return torch.from_numpy(colors.astype(np.float32) / 255.0).to("cuda") + +def pixel_art_effect(image, pixel_size, color_depth, use_aspect_ratio, palette, mask=None): + image = torch.tensor(np.array(image)).float().to("cuda") / 255.0 + height, width = image.shape[0], image.shape[1] + + if use_aspect_ratio: + aspect_ratio = width / height + pixel_size_x, pixel_size_y = pixel_size, pixel_size / aspect_ratio + else: + pixel_size_x = pixel_size_y = pixel_size + + new_width = int(width / pixel_size_x) + new_height = int(height / pixel_size_y) + + # Resize the image to create the pixelated effect + pixelated = image.permute(2, 0, 1).unsqueeze(0) + pixelated = torch.nn.functional.interpolate(pixelated, size=(new_height, new_width), mode='nearest') + pixelated = torch.nn.functional.interpolate(pixelated, size=(height, width), mode='nearest') + pixelated = pixelated.squeeze(0).permute(1, 2, 0) + + # Apply color depth reduction + pixelated = adjust_color(pixelated, color_depth) + + if palette is not None: + pixelated = apply_palette(pixelated, palette) + + if mask is not None: + mask_tensor = torch.tensor(np.array(mask)).float().to("cuda") / 255.0 + mask_tensor = mask_tensor.unsqueeze(-1).expand(-1, -1, 3) + pixelated = pixelated * mask_tensor + image * (1 - mask_tensor) + + return Image.fromarray((pixelated.cpu().numpy() * 255).astype(np.uint8)) + +def adjust_color(color, color_depth): + return torch.floor(color * color_depth) / color_depth + +def apply_palette(image, palette): + original_shape = image.shape + pixels = image.reshape(-1, 3) + distances = torch.cdist(pixels, palette) + nearest_palette_indices = torch.argmin(distances, dim=1) + new_pixels = palette[nearest_palette_indices] return new_pixels.reshape(original_shape) \ No newline at end of file diff --git a/nodes/FL_UpscaleModel.py b/nodes/FL_UpscaleModel.py new file mode 100644 index 0000000..e205611 --- /dev/null +++ b/nodes/FL_UpscaleModel.py @@ -0,0 +1,74 @@ +import torch +import comfy +from comfy_extras.nodes_upscale_model import ImageUpscaleWithModel + +class FL_UpscaleModel: + rescale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"] + precision_options = ["16", "32"] # Removed "8" as it's not standard + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "upscale" + CATEGORY = "🏵️Fill Nodes/Loaders" + + def __init__(self): + self.__imageScaler = ImageUpscaleWithModel() + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "upscale_model": ("UPSCALE_MODEL",), + "image": ("IMAGE",), + "downscale_by": ("FLOAT", { + "default": 1.0, + "min": 0.25, + "max": 1.0, + "step": 0.05, + }), + "rescale_method": (cls.rescale_methods,), + "precision": (cls.precision_options,), + } + } + + def upscale(self, upscale_model, image, downscale_by, rescale_method, precision): + original_device = image.device + original_dtype = image.dtype + + if precision == "16": + dtype = torch.float16 + else: + dtype = torch.float32 + + upscale_model = upscale_model.to(dtype).to(original_device) + image = image.to(dtype) + + with torch.no_grad(): + if dtype == torch.float16: + with torch.autocast(device_type=original_device.type, dtype=dtype): + upscaled = self.__imageScaler.upscale(upscale_model, image)[0] + else: + upscaled = self.__imageScaler.upscale(upscale_model, image)[0] + + if downscale_by < 1.0: + target_height = round(upscaled.shape[1] * downscale_by) + target_width = round(upscaled.shape[2] * downscale_by) + + # upscaled is already in [B, H, W, C] format + # We need to change it to [B, C, H, W] for interpolate + upscaled = upscaled.permute(0, 3, 1, 2) + + upscaled = torch.nn.functional.interpolate( + upscaled, + size=(target_height, target_width), + mode=rescale_method if rescale_method != "lanczos" else "bicubic", + align_corners=False if rescale_method in ["bilinear", "bicubic"] else None + ) + + # Change back to [B, H, W, C] + upscaled = upscaled.permute(0, 2, 3, 1) + + # Only clamp and convert if necessary + if dtype != original_dtype or downscale_by < 1.0: + upscaled = upscaled.clamp(0, 1).to(original_dtype).to(original_device) + + return (upscaled,) \ No newline at end of file