adds ChromaticAbberation and Vignette
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user