From 1ae3bbc89ae6e0d2e8c61122485bd0df837e17c2 Mon Sep 17 00:00:00 2001 From: melMass Date: Sat, 3 Jun 2023 20:05:19 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20=E2=9A=A1=EF=B8=8F=20initial=20commit?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 3 + __init__.py | 27 ++++ nodes/crop.py | 142 +++++++++++++++++ nodes/graph_utils.py | 41 +++++ nodes/image_processing.py | 322 ++++++++++++++++++++++++++++++++++++++ requirements.txt | 2 + utils.py | 14 ++ 7 files changed, 551 insertions(+) create mode 100644 .gitignore create mode 100644 __init__.py create mode 100644 nodes/crop.py create mode 100644 nodes/graph_utils.py create mode 100644 nodes/image_processing.py create mode 100644 requirements.txt create mode 100644 utils.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..38491ac --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +__pycache__ +*.py[cod] +*.onnx \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..6a0c303 --- /dev/null +++ b/__init__.py @@ -0,0 +1,27 @@ +from .nodes.image_processing import ( + ImageCompare, + Denoise, + Blur, + HSVtoRGB, + RGBtoHSV, + ColorCorrect, +) +from .nodes.crop import Crop, Uncrop, BoundingBox +from .nodes.graph_utils import IntToNumber, Modulo + + + +# NODE MAPPING +NODE_CLASS_MAPPINGS = { + "Int to Number (mtb)": IntToNumber, + "Bounding Box (mtb)": BoundingBox, + "Crop (mtb)": Crop, + "Uncrop (mtb)": Uncrop, + "ImageBlur (mtb)": Blur, + "Denoise (mtb)": Denoise, + "ImageCompare (mtb)": ImageCompare, + "RGB to HSV (mtb)": RGBtoHSV, + "HSV to RGB (mtb)": HSVtoRGB, + "Color Correct (mtb)": ColorCorrect, + "Modulo (mtb)": Modulo, +} diff --git a/nodes/crop.py b/nodes/crop.py new file mode 100644 index 0000000..a20564f --- /dev/null +++ b/nodes/crop.py @@ -0,0 +1,142 @@ +import torch +from ..utils import tensor2pil, pil2tensor +from PIL import Image, ImageFilter, ImageDraw + + +class BoundingBox: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "x": ("INT", {"default": 0, "max": 10000000, "min": 0, "step": 1}), + "y": ("INT", {"default": 0, "max": 10000000, "min": 0, "step": 1}), + "width": ( + "INT", + {"default": 256, "max": 10000000, "min": 0, "step": 1}, + ), + "height": ( + "INT", + {"default": 256, "max": 10000000, "min": 0, "step": 1}, + ), + } + } + + RETURN_TYPES = ("BBOX",) + FUNCTION = "do_crop" + CATEGORY = "image/crop" + + def do_crop(self, x, y, width, height): + return (x, y, width, height) + + +class Crop: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "mask": ("MASK",), + "x": ("INT", {"default": 0, "max": 10000000, "min": 0, "step": 1}), + "y": ("INT", {"default": 0, "max": 10000000, "min": 0, "step": 1}), + "width": ( + "INT", + {"default": 256, "max": 10000000, "min": 0, "step": 1}, + ), + "height": ( + "INT", + {"default": 256, "max": 10000000, "min": 0, "step": 1}, + ), + } + } + + RETURN_TYPES = ("IMAGE", "MASK", "BBOX") + FUNCTION = "do_crop" + + CATEGORY = "image/crop" + + def do_crop(self, image: torch.Tensor, mask, x, y, width, height): + + image = image.numpy() + mask = mask.numpy() + cropped_image = image[:, y : y + height, x : x + width, :] + cropped_mask = mask[y : y + height, x : x + width] + crop_data = (x, y, width, height) + + return ( + torch.from_numpy(cropped_image), + torch.from_numpy(cropped_mask), + crop_data, + ) + + +class Uncrop: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "crop_image": ("IMAGE",), + "bbox": ("BBOX",), + "border_blending": ( + "FLOAT", + {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "do_crop" + + CATEGORY = "image/crop" + + def do_crop(self, image, crop_image, bbox, border_blending): + def inset_border(image, border_width=20, border_color=(0)): + width, height = image.size + bordered_image = Image.new(image.mode, (width, height), border_color) + bordered_image.paste(image, (0, 0)) + draw = ImageDraw.Draw(bordered_image) + draw.rectangle( + (0, 0, width - 1, height - 1), outline=border_color, width=border_width + ) + return bordered_image + + image = tensor2pil(image) + crop_img = tensor2pil(crop_image) + crop_img = crop_img.convert("RGB") + + # uncrop the image based on the bounding box + bb_x, bb_y, bb_width, bb_height = bbox + + if border_blending > 1.0: + border_blending = 1.0 + elif border_blending < 0.0: + border_blending = 0.0 + + blend_ratio = (max(crop_img.size) / 2) * float(border_blending) + + blend = image.convert("RGBA") + mask = Image.new("L", image.size, 0) + + mask_block = Image.new("L", (bb_width, bb_height), 255) + mask_block = inset_border(mask_block, int(blend_ratio / 2), (0)) + + print(bbox) + mask.paste(mask_block, (bb_x, bb_y, bb_x + bb_width, bb_y + bb_height)) + blend.paste(crop_img, (bb_x, bb_y, bb_x + bb_width, bb_y + bb_height)) + + mask = mask.filter(ImageFilter.BoxBlur(radius=blend_ratio / 4)) + mask = mask.filter(ImageFilter.GaussianBlur(radius=blend_ratio / 4)) + + blend.putalpha(mask) + image = Image.alpha_composite(image.convert("RGBA"), blend) + + return (pil2tensor(image.convert("RGB")),) diff --git a/nodes/graph_utils.py b/nodes/graph_utils.py new file mode 100644 index 0000000..42f30d9 --- /dev/null +++ b/nodes/graph_utils.py @@ -0,0 +1,41 @@ +class Modulo: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "int": ("INT", {"default": 0, "min": 0, "max": 1e9, "step": 1}), + "mod": ("INT", {"default": 0, "min": 0, "max": 1e9, "step": 1}), + } + } + + RETURN_TYPES = ("INT",) + FUNCTION = "modulo" + CATEGORY = "number" + + def modulo(self, int, mod): + + return ((int + 1) % (mod + 1),) + + +class IntToNumber: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "int": ("INT", {"default": 0, "min": 0, "max": 1e9, "step": 1}), + } + } + + RETURN_TYPES = ("NUMBER",) + FUNCTION = "int_to_number" + CATEGORY = "number" + + def int_to_number(self, int): + + return (int,) diff --git a/nodes/image_processing.py b/nodes/image_processing.py new file mode 100644 index 0000000..8dc355e --- /dev/null +++ b/nodes/image_processing.py @@ -0,0 +1,322 @@ +import torch +from skimage.filters import gaussian +from skimage.restoration import denoise_tv_chambolle +from skimage.util import compare_images +from skimage.color import rgb2hsv, hsv2rgb +import numpy as np +import torchvision.transforms.functional as F +from PIL import Image +from ..utils import tensor2pil, pil2tensor + + +class ColorCorrect: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "clamp": ([True, False], {"default": True}), + "gamma": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01}, + ), + "contrast": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01}, + ), + "exposure": ( + "FLOAT", + {"default": 0.0, "min": -5.0, "max": 5.0, "step": 0.01}, + ), + "offset": ( + "FLOAT", + {"default": 0.0, "min": -5.0, "max": 5.0, "step": 0.01}, + ), + "hue": ( + "FLOAT", + {"default": 0.0, "min": -0.5, "max": 0.5, "step": 0.01}, + ), + "saturation": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01}, + ), + "value": ( + "FLOAT", + {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.01}, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "correct" + CATEGORY = "image/postprocessing" + + @staticmethod + def gamma_correction_tensor(image, gamma): + gamma_inv = 1.0 / gamma + return image.pow(gamma_inv) + + @staticmethod + def contrast_adjustment_tensor(image, contrast): + contrasted = (image - 0.5) * contrast + 0.5 + return torch.clamp(contrasted, 0.0, 1.0) + + @staticmethod + def exposure_adjustment_tensor(image, exposure): + return image * (2.0**exposure) + + @staticmethod + def offset_adjustment_tensor(image, offset): + return image + offset + + @staticmethod + def hsv_adjustment(image: torch.Tensor, hue, saturation, value): + image = tensor2pil(image) + hsv_image = image.convert("HSV") + + h, s, v = hsv_image.split() + + h = h.point(lambda x: (x + hue * 255) % 256) + s = s.point(lambda x: int(x * saturation)) + v = v.point(lambda x: int(x * value)) + + hsv_image = Image.merge("HSV", (h, s, v)) + rgb_image = hsv_image.convert("RGB") + + return pil2tensor(rgb_image) + + @staticmethod + def hsv_adjustment_tensor_not_working(image: torch.Tensor, hue, saturation, value): + """Abandonning for now""" + image = image.squeeze(0).permute(2, 0, 1) + + max_val, _ = image.max(dim=0, keepdim=True) + min_val, _ = image.min(dim=0, keepdim=True) + delta = max_val - min_val + + hue_image = torch.zeros_like(max_val) + mask = delta != 0.0 + + r, g, b = image[0], image[1], image[2] + hue_image[mask & (max_val == r)] = ((g - b) / delta)[ + mask & (max_val == r) + ] % 6.0 + hue_image[mask & (max_val == g)] = ((b - r) / delta)[ + mask & (max_val == g) + ] + 2.0 + hue_image[mask & (max_val == b)] = ((r - g) / delta)[ + mask & (max_val == b) + ] + 4.0 + + saturation_image = delta / (max_val + 1e-7) + value_image = max_val + + hue_image = (hue_image + hue) % 1.0 + saturation_image = torch.where( + mask, saturation * saturation_image, saturation_image + ) + value_image = value * value_image + + c = value_image * saturation_image + x = c * (1 - torch.abs((hue_image % 2) - 1)) + m = value_image - c + + prime_image = torch.zeros_like(image) + prime_image[0] = torch.where( + max_val == r, c, torch.where(max_val == g, x, prime_image[0]) + ) + prime_image[1] = torch.where( + max_val == r, x, torch.where(max_val == g, c, prime_image[1]) + ) + prime_image[2] = torch.where( + max_val == g, x, torch.where(max_val == b, c, prime_image[2]) + ) + + rgb_image = prime_image + m + + rgb_image = rgb_image.permute(1, 2, 0).unsqueeze(0) + + return rgb_image + + def correct( + self, + image: torch.Tensor, + clamp: bool, + gamma: float = 1.0, + contrast: float = 1.0, + exposure: float = 0.0, + offset: float = 0.0, + hue: float = 0.0, + saturation: float = 1.0, + value: float = 1.0, + ): + + # Apply color correction operations + image = self.gamma_correction_tensor(image, gamma) + image = self.contrast_adjustment_tensor(image, contrast) + image = self.exposure_adjustment_tensor(image, exposure) + image = self.offset_adjustment_tensor(image, offset) + image = self.hsv_adjustment(image, hue, saturation, value) + + if clamp: + image = torch.clamp(image, 0.0, 1.0) + + return (image,) + + +class HSVtoRGB: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "convert" + CATEGORY = "image/postprocessing" + + def convert(self, image): + image = image.numpy() + + image = image.squeeze() + # image = image.transpose(1,2,3,0) + image = hsv2rgb(image) + image = np.expand_dims(image, axis=0) + + # image = image.transpose(3,0,1,2) + return (torch.from_numpy(image),) + + +class RGBtoHSV: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "convert" + CATEGORY = "image/postprocessing" + + def convert(self, image): + image = image.numpy() + + # image = image.transpose(1,2,3,0) + image = np.squeeze(image) + image = rgb2hsv(image) + image = np.expand_dims(image, axis=0) + + # image = image.transpose(3,0,1,2) + return (torch.from_numpy(image),) + + +class ImageCompare: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "imageA": ("IMAGE",), + "imageB": ("IMAGE",), + "mode": ( + ["checkerboard", "diff", "blend"], + {"default": "checkerboard"}, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "compare" + CATEGORY = "image" + + def compare(self, imageA: torch.Tensor, imageB: torch.Tensor, mode): + imageA = imageA.numpy() + imageB = imageB.numpy() + + imageA = imageA.squeeze() + imageB = imageB.squeeze() + + image = compare_images(imageA, imageB, method=mode) + + image = np.expand_dims(image, axis=0) + return (torch.from_numpy(image),) + + +class Denoise: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "weight": ( + "FLOAT", + {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "denoise" + CATEGORY = "image/postprocessing" + + def denoise(self, image: torch.Tensor, weight): + image = image.numpy() + # image = image.transpose(1,2,3,0) + image = image.squeeze() + image = denoise_tv_chambolle(image, weight=weight) + + # image = image.transpose(3,0,1,2) + image = np.expand_dims(image, axis=0) + return (torch.from_numpy(image),) + + +class Blur: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "sigmaX": ( + "FLOAT", + {"default": 3.0, "min": 0.0, "max": 10.0, "step": 0.01}, + ), + "sigmaY": ( + "FLOAT", + {"default": 3.0, "min": 0.0, "max": 10.0, "step": 0.01}, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "blur" + CATEGORY = "image/postprocessing" + + def blur(self, image: torch.Tensor, sigmaX, sigmaY): + image = image.numpy() + image = image.transpose(1, 2, 3, 0) + # image = ndimage.gaussian_filter(image, sigma) + image = gaussian(image, sigma=(sigmaX, sigmaY, 0, 0)) + # (image, sigma=sigma, multichannel=True) + image = image.transpose(3, 0, 1, 2) + return (torch.from_numpy(image),) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..da6cc83 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +onnxruntime +imageio \ No newline at end of file diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..0c42c8e --- /dev/null +++ b/utils.py @@ -0,0 +1,14 @@ +from PIL import Image +import numpy as np +import torch +# Tensor to PIL (grabbed from WAS Suite) +def tensor2pil(image: torch.Tensor) -> Image.Image: + return Image.fromarray( + np.clip(255.0 * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8) + ) + + +# Convert PIL to Tensor (grabbed from WAS Suite) +def pil2tensor(image: Image.Image) -> torch.Tensor: + return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0) +