diff --git a/README.md b/README.md index 6406dc6..8709d31 100644 --- a/README.md +++ b/README.md @@ -17,6 +17,7 @@ Both images have the workflow attached, and it is included so feel free to use i - Blend: Blends two images together with a variety of different modes - Blur: Applies a Gaussian blur to the input image, softening the details - CannyEdgeDetection: Applies Canny edge detection to the input image + - Chromatic Aberration: Shifts the color channels in an image, creating a glitch aesthetic - ColorCorrect: Adjusts the color balance, temperature, hue, brightness, contrast, saturation, and gamma of an image - Dissolve: Creates a grainy blend of two images using random pixels based on a dissolve factor. - DodgeAndBurn: Adjusts image brightness using dodge and burn effects based on a mask and intensity. @@ -28,6 +29,7 @@ Both images have the workflow attached, and it is included so feel free to use i - Quantize: Set and dither the amount of colors in an image from 0-256, reducing color information - Sharpen: Enhances the details in an image by applying a sharpening filter - Solarize: Inverts image colors based on a threshold for a striking, high-contrast effect + - Vignette: Applies a vignette effect, putting the corners of the image in shadow ## Combine Nodes diff --git a/post_processing/chromatic_abberation.py b/post_processing/chromatic_abberation.py new file mode 100644 index 0000000..e291416 --- /dev/null +++ b/post_processing/chromatic_abberation.py @@ -0,0 +1,57 @@ +import torch + +class ChromaticAberration: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "red_shift": ("INT", { + "default": 0, + "min": -20, + "max": 20, + "step": 1 + }), + "red_direction": (["horizontal", "vertical"],), + "green_shift": ("INT", { + "default": 0, + "min": -20, + "max": 20, + "step": 1 + }), + "green_direction": (["horizontal", "vertical"],), + "blue_shift": ("INT", { + "default": 0, + "min": -20, + "max": 20, + "step": 1 + }), + "blue_direction": (["horizontal", "vertical"],), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "chromatic_aberration" + + CATEGORY = "postprocessing" + + def chromatic_aberration(self, image: torch.Tensor, red_shift: int, green_shift: int, blue_shift: int, red_direction: str, green_direction: str, blue_direction: str): + def get_shift(direction, shift): + shift = -shift if direction == 'vertical' else shift # invert vertical shift as otherwise positive actually shifts down + return (shift, 0) if direction == 'vertical' else (0, shift) + + x = image.permute(0, 3, 1, 2) + shifts = [get_shift(direction, shift) for direction, shift in zip([red_direction, green_direction, blue_direction], [red_shift, green_shift, blue_shift])] + channels = [torch.roll(x[:, i, :, :], shifts=shifts[i], dims=(1, 2)) for i in range(3)] + + output = torch.stack(channels, dim=1) + output = output.permute(0, 2, 3, 1) + + return (output,) + +NODE_CLASS_MAPPINGS = { + "ChromaticAberration": ChromaticAberration +} diff --git a/post_processing/vignette.py b/post_processing/vignette.py new file mode 100644 index 0000000..10f9b28 --- /dev/null +++ b/post_processing/vignette.py @@ -0,0 +1,47 @@ +import numpy as np +import torch + +class Vignette: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "a": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 10.0, + "step": 1.0 + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "apply_vignette" + + CATEGORY = "postprocessing" + + def apply_vignette(self, image: torch.Tensor, vignette: float): + if vignette == 0: + return (image,) + height, width, _ = image.shape[-3:] + x = torch.linspace(-1, 1, width, device=image.device) + y = torch.linspace(-1, 1, height, device=image.device) + X, Y = torch.meshgrid(x, y, indexing="ij") + radius = torch.sqrt(X ** 2 + Y ** 2) + + # Map vignette strength from 0-10 to 1.800-0.800 + mapped_vignette_strength = 1.8 - (vignette - 1) * 0.1 + vignette = 1 - torch.clamp(radius / mapped_vignette_strength, 0, 1) + vignette = vignette[..., None] + + vignette_image = torch.clamp(image * vignette, 0, 1) + + return (vignette_image,) + +NODE_CLASS_MAPPINGS = { + "Vignette": Vignette, +} diff --git a/post_processing_nodes.py b/post_processing_nodes.py index 054ff84..2f8ce25 100644 --- a/post_processing_nodes.py +++ b/post_processing_nodes.py @@ -4,6 +4,12 @@ import cv2 import numpy as np from PIL import Image, ImageEnhance from PIL import Image +import imageio.v2 as imageio +from math import sqrt +import sys +import argparse +import os +import random class ArithmeticBlend: @@ -207,6 +213,58 @@ class CannyEdgeDetection: return (result,) +class ChromaticAberration: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "red_shift": ("INT", { + "default": 0, + "min": -20, + "max": 20, + "step": 1 + }), + "red_direction": (["horizontal", "vertical"],), + "green_shift": ("INT", { + "default": 0, + "min": -20, + "max": 20, + "step": 1 + }), + "green_direction": (["horizontal", "vertical"],), + "blue_shift": ("INT", { + "default": 0, + "min": -20, + "max": 20, + "step": 1 + }), + "blue_direction": (["horizontal", "vertical"],), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "chromatic_aberration" + + CATEGORY = "postprocessing" + + def chromatic_aberration(self, image: torch.Tensor, red_shift: int, green_shift: int, blue_shift: int, red_direction: str, green_direction: str, blue_direction: str): + def get_shift(direction, shift): + shift = -shift if direction == 'vertical' else shift # invert vertical shift as otherwise positive actually shifts down + return (shift, 0) if direction == 'vertical' else (0, shift) + + x = image.permute(0, 3, 1, 2) + shifts = [get_shift(direction, shift) for direction, shift in zip([red_direction, green_direction, blue_direction], [red_shift, green_shift, blue_shift])] + channels = [torch.roll(x[:, i, :, :], shifts=shifts[i], dims=(1, 2)) for i in range(3)] + + output = torch.stack(channels, dim=1) + output = output.permute(0, 2, 3, 1) + + return (output,) + class ColorCorrect: def __init__(self): pass @@ -893,6 +951,123 @@ class Solarize: solarized_image = torch.clamp(solarized_image, 0, 1) return (solarized_image,) +class Vignette: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "a": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 10.0, + "step": 1.0 + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "apply_vignette" + + CATEGORY = "postprocessing" + + def apply_vignette(self, image: torch.Tensor, vignette: float): + if vignette == 0: + return (image,) + height, width, _ = image.shape[-3:] + x = torch.linspace(-1, 1, width, device=image.device) + y = torch.linspace(-1, 1, height, device=image.device) + X, Y = torch.meshgrid(x, y, indexing="ij") + radius = torch.sqrt(X ** 2 + Y ** 2) + + # Map vignette strength from 0-10 to 1.800-0.800 + mapped_vignette_strength = 1.8 - (vignette - 1) * 0.1 + vignette = 1 - torch.clamp(radius / mapped_vignette_strength, 0, 1) + vignette = vignette[..., None] + + vignette_image = torch.clamp(image * vignette, 0, 1) + + return (vignette_image,) + +class ElectroShock: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "glow_intensity": ("INT", {"default": 50, "min": 0, "max": 100, "step": 1}), + "line_frequency": ("INT", {"default": 25, "min": 0, "max": 100, "step": 1}), + "line_thickness": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}), + "random_seed": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1}), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "electro_shock" + + CATEGORY = "effects" + + def midpoint_displacement(self, x1, y1, x2, y2, displacement, mask, line_thickness): + if abs(x2 - x1) < 2 and abs(y2 - y1) < 2: + return + + mid_x = (x1 + x2) // 2 + mid_y = (y1 + y2) // 2 + + mid_x += int(random.uniform(-displacement, displacement)) + mid_y += int(random.uniform(-displacement, displacement)) + + cv2.line(mask, (x1, y1), (mid_x, mid_y), 255, line_thickness) + cv2.line(mask, (mid_x, mid_y), (x2, y2), 255, line_thickness) + + self.midpoint_displacement(x1, y1, mid_x, mid_y, displacement / 2, mask, line_thickness) + self.midpoint_displacement(mid_x, mid_y, x2, y2, displacement / 2, mask, line_thickness) + + + def electro_shock(self, image: torch.Tensor, glow_intensity: int, line_frequency: int, line_thickness: int, random_seed: int = None): + if random_seed is not None: + random.seed(random_seed) + np.random.seed(random_seed) + + line_color = [255, 255, 255] + + batch_size, height, width, _ = image.shape + result = torch.zeros_like(image) + + for b in range(batch_size): + tensor_image = image[b] + img = (tensor_image * 255).to(torch.uint8).numpy() + + # Apply the ElectroShock effect using OpenCV functions + mask = np.zeros((height, width), np.uint8) + num_lines = int(line_frequency * (height * width) / 100000) + initial_displacement = int(height / 8) + + for _ in range(num_lines): + x1, y1 = random.randint(0, width - 1), random.randint(0, height - 1) + x2, y2 = random.randint(0, width - 1), random.randint(0, height - 1) + self.midpoint_displacement(x1, y1, x2, y2, initial_displacement, mask, line_thickness) + + # Apply glow effect + glow_radius = int(glow_intensity * 0.1) + mask_blurred = cv2.GaussianBlur(mask, (glow_radius * 2 + 1, glow_radius * 2 + 1), 0) + + # Add glow to the original image + colored_mask = cv2.cvtColor(mask_blurred, cv2.COLOR_GRAY2BGR) + colored_mask[np.where((colored_mask == [255, 255, 255]).all(axis=2))] = line_color + electro_shock_img = cv2.addWeighted(img, 1, colored_mask, glow_intensity / 100, 0) + + electro_shock_array = torch.tensor(electro_shock_img).float() / 255 + result[b] = electro_shock_array + + return (result,) + def gaussian_kernel(kernel_size: int, sigma: float): x, y = torch.meshgrid(torch.linspace(-1, 1, kernel_size), torch.linspace(-1, 1, kernel_size), indexing="ij") d = torch.sqrt(x * x + y * y) @@ -990,11 +1165,122 @@ def pixel_sort(img, mask, horizontal_sort=False, span_limit=None, sort_by='H', r return sorted_image +def get_fish_xn_yn(source_x, source_y, radius, distortion): + """ + Get normalized x, y pixel coordinates from the original image and return normalized + x, y pixel coordinates in the destination fished image. + :param distortion: Amount in which to move pixels from/to center. + As distortion grows, pixels will be moved further from the center, and vice versa. + """ + + if 1 - distortion*(radius**2) == 0: + return source_x, source_y + + return source_x / (1 - (distortion*(radius**2))), source_y / (1 - (distortion*(radius**2))) + + + +def fish(img, distortion_coefficient): + """ + :type img: numpy.ndarray + :param distortion_coefficient: The amount of distortion to apply. + :return: numpy.ndarray - the image with applied effect. + """ + + # If input image is only BW or RGB convert it to RGBA + # So that output 'frame' can be transparent. + w, h = img.shape[0], img.shape[1] + if len(img.shape) == 2: + # Duplicate the one BW channel twice to create Black and White + # RGB image (For each pixel, the 3 channels have the same value) + bw_channel = np.copy(img) + img = np.dstack((img, bw_channel)) + img = np.dstack((img, bw_channel)) + if len(img.shape) == 3 and img.shape[2] == 3: + print("RGB to RGBA") + img = np.dstack((img, np.full((w, h), 255))) + + # prepare array for dst image + dstimg = np.zeros_like(img) + + # floats for calcultions + w, h = float(w), float(h) + + # easier calcultion if we traverse x, y in dst image + for x in range(len(dstimg)): + for y in range(len(dstimg[x])): + + # normalize x and y to be in interval of [-1, 1] + xnd, ynd = float((2*x - w)/w), float((2*y - h)/h) + + # get xn and yn distance from normalized center + rd = sqrt(xnd**2 + ynd**2) + + # new normalized pixel coordinates + xdu, ydu = get_fish_xn_yn(xnd, ynd, rd, distortion_coefficient) + + # convert the normalized distorted xdn and ydn back to image pixels + xu, yu = int(((xdu + 1)*w)/2), int(((ydu + 1)*h)/2) + + # if new pixel is in bounds copy from source pixel to destination pixel + if 0 <= xu and xu < img.shape[0] and 0 <= yu and yu < img.shape[1]: + dstimg[x][y] = img[xu][yu] + + return dstimg.astype(np.uint8) + + + +def parse_args(args=sys.argv[1:]): + """Parse arguments.""" + + parser = argparse.ArgumentParser( + description="Apply fish-eye effect to images.", + prog='python3 fish.py') + + parser.add_argument("-i", "--image", help="path to image file." + " If no input is given, the supplied example 'grid.jpg' will be used.", + type=str, default="test.png") + + parser.add_argument("-o", "--outpath", help="file path to write output to." + " format: .", + type=str, default="fish.png") + + parser.add_argument("-d", "--distortion", + help="The distoration coefficient. How much the move pixels from/to the center." + " Recommended values are between -1 and 1." + " The bigger the distortion, the further pixels will be moved outwars from the center (fisheye)." + " The Smaller the distortion, the closer pixels will be move inwards toward the center (rectilinear)." + " For example, to reverse the fisheye effect with --distoration 0.5," + " You can run with --distortion -0.3." + " Note that due to double processing the result will be somewhat distorted.", + type=float, default=0.5) + + return parser.parse_args(args) + + +if __name__ == "__main__": + args = parse_args() + try: + imgobj = imageio.imread(args.image) + except Exception as e: + print(e) + sys.exit(1) + if os.path.exists(args.outpath): + ans = input( + args.outpath + " exists. File will be overridden. Continue? y/n: ") + if ans.lower() != 'y': + print("exiting") + sys.exit(0) + + output_img = fish(imgobj, args.distortion) + imageio.imwrite(args.outpath, output_img, format='png') + NODE_CLASS_MAPPINGS = { "ArithmeticBlend": ArithmeticBlend, "Blend": Blend, "Blur": Blur, "CannyEdgeDetection": CannyEdgeDetection, + "ChromaticAberration": ChromaticAberration, "ColorCorrect": ColorCorrect, "Dissolve": Dissolve, "DodgeAndBurn": DodgeAndBurn, @@ -1006,4 +1292,6 @@ NODE_CLASS_MAPPINGS = { "Quantize": Quantize, "Sharpen": Sharpen, "Solarize": Solarize, + "Vignette": Vignette, + "ElectroShock": ElectroShock, }