Upscale Model node added

This commit is contained in:
Fillip
2024-08-14 15:54:34 -07:00
parent b9ab9f1373
commit 462dbc9bb3
3 changed files with 200 additions and 123 deletions
+4 -1
View File
@@ -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",
}
+122 -122
View File
@@ -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)
+74
View File
@@ -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,)