diff --git a/README.md b/README.md index 79cfeda..6c2ac87 100644 --- a/README.md +++ b/README.md @@ -1,9 +1,29 @@ # ComfyUI-post-processing-nodes -A collection of post processing nodes for [ComfyUI](https://github.com/comfyanonymous/ComfyUI), simply download this repo and drag `combined_nodes.py` into your `custom_nodes/` folder +A collection of post processing nodes for [ComfyUI](https://github.com/comfyanonymous/ComfyUI), simply download this repo and drag `post_processing_nodes.py` into your `custom_nodes/` folder + +## Node List + + - Blend: Blends two images together with a variety of different modes + - CannyEdgeDetection: Applies Canny edge detection to the input image + - ColorCorrect: Adjusts the color balance, temperature, hue, brightness, contrast, saturation, and gamma of an image + - Dither: Reduces the color information in an image by dithering, resulting in a patterned, pixelated appearance + - FilmGrain: Adds a film grain effect to the image, along with options to control the temperature, and vignetting + - GaussianBlur: Applies a Gaussian blur to the input image, softening the details + - KMeansQuantize: Reduce the amount of colors in an image from 0-256 + - PixelSort: Rearranges the pixels in the input image based on their values, and input mask. Creates a cool glitch like effect. + - Sharpen: Enhances the details in an image by applying a sharpening filter + +## Example workflow + +![__image__](images/example-workflow.png) ## Combine Nodes -By default `combined_nodes.py` should have all of the combined nodes. If you want a subset of nodes, you can run +By default `post_processing_nodes.py` should have all of the combined nodes. If you want a subset of nodes, you can run python combine_files.py [--files FILES [FILES ...]] [--output OUTPUT] + +or just run + + python combine_files.py -h for more help \ No newline at end of file diff --git a/combine_files.py b/combine_files.py index 7b3727d..f695913 100644 --- a/combine_files.py +++ b/combine_files.py @@ -1,3 +1,4 @@ +from collections import OrderedDict from pathlib import Path import sys import os @@ -6,16 +7,17 @@ import ast import argparse -def get_python_files(path): - for file in Path(path).glob("*.py"): - if file.is_file() and not file.name.startswith("combine"): - yield str(file) + +def get_python_files(path, recursive=False, args=None): + search_pattern = "**/*.py" if recursive else "*.py" + files = sorted([str(file) for file in Path(path).glob(search_pattern) if file.is_file() and not file.name.startswith("combine") and not args.output in str(file)]) + yield from files def parse_files(files): - imports = set() - class_definitions = set() - node_class_mappings = set() - functions = set() + imports = OrderedDict() + class_definitions = OrderedDict() + node_class_mappings = OrderedDict() + functions = OrderedDict() for file in files: # read file as lines @@ -30,7 +32,7 @@ def parse_files(files): while i < num_lines: line = lines[i] if line.startswith("import") or line.startswith("from"): - imports.add(line.strip()) + imports[line.strip()] = None elif line.startswith("class"): class_info = line @@ -38,11 +40,11 @@ def parse_files(files): while not lines[j].startswith("NODE_CLASS_MAPPINGS"): class_info += lines[j] j += 1 - class_definitions.add(class_info) + class_definitions[class_info] = None i = j - 1 elif line.startswith("NODE_CLASS_MAPPINGS"): - node_class_mappings.add(lines[i+1]) + node_class_mappings[lines[i+1]] = None elif line.startswith("def"): function_info = line @@ -50,7 +52,7 @@ def parse_files(files): while j < num_lines and not lines[j].startswith("NODE_CLASS_MAPPINGS") and not lines[j].startswith("def"): function_info += lines[j] j += 1 - functions.add(function_info) + functions[function_info] = None i = j - 1 i += 1 @@ -91,12 +93,19 @@ def main(): parser = argparse.ArgumentParser(description="Collect unique imports from Python files") parser.add_argument("--all", action="store_true", help="Include all Python files in the specified directory") parser.add_argument("--files", nargs="+", help="Specify Python files to parse") - parser.add_argument("--output", default="combined_nodes.py", help="Specify the output file name") + parser.add_argument("--folder", default=".", help="Specify a folder to search for files") + parser.add_argument("--output", default="post_processing_nodes.py", help="Specify the output file name") args = parser.parse_args() + args.all = True + if args.all: - args.path = "." - files = get_python_files(args.path) + args.folder = "." if args.folder is None else args.folder + files = get_python_files(args.folder, recursive=True, args=args) + imports, class_definitions, node_class_mappings, functions = parse_files(files) + write_combined(imports, class_definitions, node_class_mappings, functions, args.output) + elif args.folder is not None: + files = get_python_files(args.folder, recursive=True, args=args) imports, class_definitions, node_class_mappings, functions = parse_files(files) write_combined(imports, class_definitions, node_class_mappings, functions, args.output) else: diff --git a/images/example-workflow.png b/images/example-workflow.png new file mode 100644 index 0000000..7ef08b5 Binary files /dev/null and b/images/example-workflow.png differ diff --git a/blend.py b/post_processing/blend.py similarity index 100% rename from blend.py rename to post_processing/blend.py diff --git a/canny_edge_detect.py b/post_processing/canny_edge_detect.py similarity index 100% rename from canny_edge_detect.py rename to post_processing/canny_edge_detect.py diff --git a/color_correct.py b/post_processing/color_correct.py similarity index 100% rename from color_correct.py rename to post_processing/color_correct.py diff --git a/dither.py b/post_processing/dither.py similarity index 100% rename from dither.py rename to post_processing/dither.py diff --git a/film_grain.py b/post_processing/film_grain.py similarity index 100% rename from film_grain.py rename to post_processing/film_grain.py diff --git a/gaussian_blur.py b/post_processing/gaussian_blur.py similarity index 100% rename from gaussian_blur.py rename to post_processing/gaussian_blur.py diff --git a/kmeans_quantize.py b/post_processing/kmeans_quantize.py similarity index 100% rename from kmeans_quantize.py rename to post_processing/kmeans_quantize.py diff --git a/pixel_sort.py b/post_processing/pixel_sort.py similarity index 100% rename from pixel_sort.py rename to post_processing/pixel_sort.py diff --git a/sharpen.py b/post_processing/sharpen.py similarity index 100% rename from sharpen.py rename to post_processing/sharpen.py diff --git a/combined_nodes.py b/post_processing_nodes.py similarity index 96% rename from combined_nodes.py rename to post_processing_nodes.py index 35ef8b0..2a50096 100644 --- a/combined_nodes.py +++ b/post_processing_nodes.py @@ -1,10 +1,58 @@ -import numpy as np import torch import cv2 -import torch.nn.functional as F +import numpy as np from PIL import Image, ImageEnhance +import torch.nn.functional as F +class Blend: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image1": ("IMAGE",), + "image2": ("IMAGE",), + "blend_factor": ("FLOAT", { + "default": 0.5, + "min": 0.0, + "max": 1.0, + "step": 0.01 + }), + "blend_mode": (["normal", "multiply", "screen", "overlay", "soft_light"],), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "blend_images" + + CATEGORY = "postprocessing" + + def blend_images(self, image1: torch.Tensor, image2: torch.Tensor, blend_factor: float, blend_mode: str): + blended_image = self.blend_mode(image1, image2, blend_mode) + blended_image = image1 * (1 - blend_factor) + blended_image * blend_factor + blended_image = torch.clamp(blended_image, 0, 1) + return (blended_image,) + + def blend_mode(self, img1, img2, mode): + if mode == "normal": + return img2 + elif mode == "multiply": + return img1 * img2 + elif mode == "screen": + return 1 - (1 - img1) * (1 - img2) + elif mode == "overlay": + return torch.where(img1 <= 0.5, 2 * img1 * img2, 1 - 2 * (1 - img1) * (1 - img2)) + elif mode == "soft_light": + return torch.where(img2 <= 0.5, img1 - (1 - 2 * img2) * img1 * (1 - img1), img1 + (2 * img2 - 1) * (self.g(img1) - img1)) + else: + raise ValueError(f"Unsupported blend mode: {mode}") + + def g(self, x): + return torch.where(x <= 0.25, ((16 * x - 12) * x + 4) * x, torch.sqrt(x)) + class CannyEdgeDetection: def __init__(self): pass @@ -47,6 +95,112 @@ class CannyEdgeDetection: return (result,) +class ColorCorrect: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "temperature": ("FLOAT", { + "default": 0, + "min": -100, + "max": 100, + "step": 5 + }), + "hue": ("FLOAT", { + "default": 0, + "min": -90, + "max": 90, + "step": 5 + }), + "brightness": ("FLOAT", { + "default": 0, + "min": -100, + "max": 100, + "step": 5 + }), + "contrast": ("FLOAT", { + "default": 0, + "min": -100, + "max": 100, + "step": 5 + }), + "saturation": ("FLOAT", { + "default": 0, + "min": -100, + "max": 100, + "step": 5 + }), + "gamma": ("FLOAT", { + "default": 1, + "min": 0.2, + "max": 2.2, + "step": 0.1 + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "color_correct" + + CATEGORY = "postprocessing" + + def color_correct(self, image: torch.Tensor, temperature: float, hue: float, brightness: float, contrast: float, saturation: float, gamma: float): + batch_size, height, width, _ = image.shape + result = torch.zeros_like(image) + + for b in range(batch_size): + tensor_image = image[b].numpy() + + brightness /= 100 + contrast /= 100 + saturation /= 100 + temperature /= 100 + + brightness = 1 + brightness + contrast = 1 + contrast + saturation = 1 + saturation + + modified_image = Image.fromarray((tensor_image * 255).astype(np.uint8)) + + # brightness + modified_image = ImageEnhance.Brightness(modified_image).enhance(brightness) + + # contrast + modified_image = ImageEnhance.Contrast(modified_image).enhance(contrast) + modified_image = np.array(modified_image).astype(np.float32) + + # temperature + if temperature > 0: + modified_image[:, :, 0] *= 1 + temperature + modified_image[:, :, 1] *= 1 + temperature * 0.4 + elif temperature < 0: + modified_image[:, :, 2] *= 1 - temperature + modified_image = np.clip(modified_image, 0, 255)/255 + + # gamma + modified_image = np.clip(np.power(modified_image, gamma), 0, 1) + + # saturation + hls_img = cv2.cvtColor(modified_image, cv2.COLOR_RGB2HLS) + hls_img[:, :, 2] = np.clip(saturation*hls_img[:, :, 2], 0, 1) + modified_image = cv2.cvtColor(hls_img, cv2.COLOR_HLS2RGB) * 255 + + # hue + hsv_img = cv2.cvtColor(modified_image, cv2.COLOR_RGB2HSV) + hsv_img[:, :, 0] = (hsv_img[:, :, 0] + hue) % 360 + modified_image = cv2.cvtColor(hsv_img, cv2.COLOR_HSV2RGB) + + modified_image = modified_image.astype(np.uint8) + modified_image = modified_image / 255 + modified_image = torch.from_numpy(modified_image).unsqueeze(0) + result[b] = modified_image + + return (result, ) + class Dither: def __init__(self): pass @@ -295,6 +449,62 @@ class GaussianBlur: return (blurred,) +class KMeansQuantize: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "colors": ("INT", { + "default": 16, + "min": 1, + "max": 256, + "step": 1 + }), + "precision": ("INT", { + "default": 10, + "min": 1, + "max": 100, + "step": 1 + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "kmeans_quantize" + + CATEGORY = "postprocessing" + + def kmeans_quantize(self, image: torch.Tensor, colors: int, precision: int): + batch_size, height, width, _ = image.shape + result = torch.zeros_like(image) + + for b in range(batch_size): + tensor_image = image[b].numpy().astype(np.float32) + img = tensor_image + + height, width, c = img.shape + + criteria = ( + cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, + precision * 5, 0.01 + ) + + img_copy = img.reshape(-1, c) + _, label, center = cv2.kmeans( + img_copy, colors, None, + criteria, 1, cv2.KMEANS_PP_CENTERS + ) + + img = center[label.flatten()].reshape(*img.shape) + tensor = torch.from_numpy(img).unsqueeze(0) + result[b] = tensor + + return (result,) + class PixelSort: def __init__(self): pass @@ -385,215 +595,36 @@ class Sharpen: return (result,) -class ColorCorrect: - def __init__(self): - pass +def sort_span(span, sort_by, reverse_sorting): + if sort_by == 'H': + key = lambda x: x[1][0] + elif sort_by == 'S': + key = lambda x: x[1][1] + else: + key = lambda x: x[1][2] - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "image": ("IMAGE",), - "temperature": ("FLOAT", { - "default": 0, - "min": -100, - "max": 100, - "step": 5 - }), - "hue": ("FLOAT", { - "default": 0, - "min": -90, - "max": 90, - "step": 5 - }), - "brightness": ("FLOAT", { - "default": 0, - "min": -100, - "max": 100, - "step": 5 - }), - "contrast": ("FLOAT", { - "default": 0, - "min": -100, - "max": 100, - "step": 5 - }), - "saturation": ("FLOAT", { - "default": 0, - "min": -100, - "max": 100, - "step": 5 - }), - "gamma": ("FLOAT", { - "default": 1, - "min": 0.2, - "max": 2.2, - "step": 0.1 - }), - }, - } + span = sorted(span, key=key, reverse=reverse_sorting) + return [x[0] for x in span] - RETURN_TYPES = ("IMAGE",) - FUNCTION = "color_correct" - CATEGORY = "postprocessing" +def find_spans(mask, span_limit=None): + spans = [] + start = None + for i, value in enumerate(mask): + if value == 0 and start is None: + start = i + if value == 1 and start is not None: + span_length = i - start + if span_limit is None or span_length <= span_limit: + spans.append((start, i)) + start = None + if start is not None: + span_length = len(mask) - start + if span_limit is None or span_length <= span_limit: + spans.append((start, len(mask))) - def color_correct(self, image: torch.Tensor, temperature: float, hue: float, brightness: float, contrast: float, saturation: float, gamma: float): - batch_size, height, width, _ = image.shape - result = torch.zeros_like(image) + return spans - for b in range(batch_size): - tensor_image = image[b].numpy() - - brightness /= 100 - contrast /= 100 - saturation /= 100 - temperature /= 100 - - brightness = 1 + brightness - contrast = 1 + contrast - saturation = 1 + saturation - - modified_image = Image.fromarray((tensor_image * 255).astype(np.uint8)) - - # brightness - modified_image = ImageEnhance.Brightness(modified_image).enhance(brightness) - - # contrast - modified_image = ImageEnhance.Contrast(modified_image).enhance(contrast) - modified_image = np.array(modified_image).astype(np.float32) - - # temperature - if temperature > 0: - modified_image[:, :, 0] *= 1 + temperature - modified_image[:, :, 1] *= 1 + temperature * 0.4 - elif temperature < 0: - modified_image[:, :, 2] *= 1 - temperature - modified_image = np.clip(modified_image, 0, 255)/255 - - # gamma - modified_image = np.clip(np.power(modified_image, gamma), 0, 1) - - # saturation - hls_img = cv2.cvtColor(modified_image, cv2.COLOR_RGB2HLS) - hls_img[:, :, 2] = np.clip(saturation*hls_img[:, :, 2], 0, 1) - modified_image = cv2.cvtColor(hls_img, cv2.COLOR_HLS2RGB) * 255 - - # hue - hsv_img = cv2.cvtColor(modified_image, cv2.COLOR_RGB2HSV) - hsv_img[:, :, 0] = (hsv_img[:, :, 0] + hue) % 360 - modified_image = cv2.cvtColor(hsv_img, cv2.COLOR_HSV2RGB) - - modified_image = modified_image.astype(np.uint8) - modified_image = modified_image / 255 - modified_image = torch.from_numpy(modified_image).unsqueeze(0) - result[b] = modified_image - - return (result, ) - -class KMeansQuantize: - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "image": ("IMAGE",), - "colors": ("INT", { - "default": 16, - "min": 1, - "max": 256, - "step": 1 - }), - "precision": ("INT", { - "default": 10, - "min": 1, - "max": 100, - "step": 1 - }), - }, - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "kmeans_quantize" - - CATEGORY = "postprocessing" - - def kmeans_quantize(self, image: torch.Tensor, colors: int, precision: int): - batch_size, height, width, _ = image.shape - result = torch.zeros_like(image) - - for b in range(batch_size): - tensor_image = image[b].numpy().astype(np.float32) - img = tensor_image - - height, width, c = img.shape - - criteria = ( - cv2.TERM_CRITERIA_EPS + cv2.TERM_CRITERIA_MAX_ITER, - precision * 5, 0.01 - ) - - img_copy = img.reshape(-1, c) - _, label, center = cv2.kmeans( - img_copy, colors, None, - criteria, 1, cv2.KMEANS_PP_CENTERS - ) - - img = center[label.flatten()].reshape(*img.shape) - tensor = torch.from_numpy(img).unsqueeze(0) - result[b] = tensor - - return (result,) - -class Blend: - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "image1": ("IMAGE",), - "image2": ("IMAGE",), - "blend_factor": ("FLOAT", { - "default": 0.5, - "min": 0.0, - "max": 1.0, - "step": 0.01 - }), - "blend_mode": (["normal", "multiply", "screen", "overlay", "soft_light"],), - }, - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "blend_images" - - CATEGORY = "postprocessing" - - def blend_images(self, image1: torch.Tensor, image2: torch.Tensor, blend_factor: float, blend_mode: str): - blended_image = self.blend_mode(image1, image2, blend_mode) - blended_image = image1 * (1 - blend_factor) + blended_image * blend_factor - blended_image = torch.clamp(blended_image, 0, 1) - return (blended_image,) - - def blend_mode(self, img1, img2, mode): - if mode == "normal": - return img2 - elif mode == "multiply": - return img1 * img2 - elif mode == "screen": - return 1 - (1 - img1) * (1 - img2) - elif mode == "overlay": - return torch.where(img1 <= 0.5, 2 * img1 * img2, 1 - 2 * (1 - img1) * (1 - img2)) - elif mode == "soft_light": - return torch.where(img2 <= 0.5, img1 - (1 - 2 * img2) * img1 * (1 - img1), img1 + (2 * img2 - 1) * (self.g(img1) - img1)) - else: - raise ValueError(f"Unsupported blend mode: {mode}") - - def g(self, x): - return torch.where(x <= 0.25, ((16 * x - 12) * x + 4) * x, torch.sqrt(x)) def pixel_sort(img, mask, horizontal_sort=False, span_limit=None, sort_by='H', reverse_sorting=False): height, width, _ = img.shape @@ -655,45 +686,14 @@ def pixel_sort(img, mask, horizontal_sort=False, span_limit=None, sort_by='H', r return sorted_image -def sort_span(span, sort_by, reverse_sorting): - if sort_by == 'H': - key = lambda x: x[1][0] - elif sort_by == 'S': - key = lambda x: x[1][1] - else: - key = lambda x: x[1][2] - - span = sorted(span, key=key, reverse=reverse_sorting) - return [x[0] for x in span] - - -def find_spans(mask, span_limit=None): - spans = [] - start = None - for i, value in enumerate(mask): - if value == 0 and start is None: - start = i - if value == 1 and start is not None: - span_length = i - start - if span_limit is None or span_length <= span_limit: - spans.append((start, i)) - start = None - if start is not None: - span_length = len(mask) - start - if span_limit is None or span_length <= span_limit: - spans.append((start, len(mask))) - - return spans - - NODE_CLASS_MAPPINGS = { "Blend": Blend, - "GaussianBlur": GaussianBlur, - "PixelSort": PixelSort, - "FilmGrain": FilmGrain, - "ColorCorrect": ColorCorrect, - "Sharpen": Sharpen, "CannyEdgeDetection": CannyEdgeDetection, - "KMeansQuantize": KMeansQuantize, + "ColorCorrect": ColorCorrect, "Dither": Dither, + "FilmGrain": FilmGrain, + "GaussianBlur": GaussianBlur, + "KMeansQuantize": KMeansQuantize, + "PixelSort": PixelSort, + "Sharpen": Sharpen, }