Upscale Model node added
This commit is contained in:
+4
-1
@@ -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",
|
||||
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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,)
|
||||
Reference in New Issue
Block a user