From 85174002137d86b42510fa76539010243efb7c04 Mon Sep 17 00:00:00 2001 From: spacepxl Date: Fri, 23 Feb 2024 04:17:49 -0500 Subject: [PATCH] RelightSimple --- nodes.py | 36 ++++++++++++++++++++++++++++++++++++ 1 file changed, 36 insertions(+) diff --git a/nodes.py b/nodes.py index 8b8296a..9e5ce53 100644 --- a/nodes.py +++ b/nodes.py @@ -1,5 +1,6 @@ import os import sys +import math import copy import torch import torchvision.transforms @@ -886,6 +887,39 @@ class OffsetLatentImage: latent[:,3,:,:] = offset_3 return ({"samples":latent}, ) +class RelightSimple: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "normals": ("IMAGE",), + "x_dir": ("FLOAT", {"default": 0.0, "min": -1.5, "max": 1.5, "step": 0.01}), + "y_dir": ("FLOAT", {"default": 0.0, "min": -1.5, "max": 1.5, "step": 0.01}), + "brightness": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100, "step": 0.01}), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "relight" + + CATEGORY = "image/filters" + + def relight(self, image, normals, x_dir, y_dir, brightness): + if image.shape[0] != normals.shape[0]: + raise Exception("Batch size for image and normals must match") + norm = normals.detach().clone() * 2 - 1 + norm = torch.nn.functional.interpolate(norm.movedim(-1,1), size=(image.shape[1], image.shape[2]), mode='bilinear').movedim(1,-1) + light = torch.tensor([x_dir, y_dir, abs(1 - math.sqrt(x_dir ** 2 + y_dir ** 2) * 0.7)]) + light = torch.nn.functional.normalize(light, dim=0) + + diffuse = norm[:,:,:,0] * light[0] + norm[:,:,:,1] * light[1] + norm[:,:,:,2] * light[2] + diffuse = torch.clip(diffuse.unsqueeze(3).repeat(1,1,1,3), 0, 1) + + relit = image.detach().clone() + relit[:,:,:,:3] = torch.clip(relit[:,:,:,:3] * diffuse * brightness, 0, 1) + return (relit,) + class LatentStats: @classmethod def INPUT_TYPES(s): @@ -1413,6 +1447,7 @@ NODE_CLASS_MAPPINGS = { "LatentStats": LatentStats, "NormalMapSimple": NormalMapSimple, "OffsetLatentImage": OffsetLatentImage, + "RelightSimple": RelightSimple, "RemapRange": RemapRange, "ShuffleChannels": ShuffleChannels, "Tonemap": Tonemap, @@ -1448,6 +1483,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LatentStats": "Latent Stats", "NormalMapSimple": "Normal Map (Simple)", "OffsetLatentImage": "Offset Latent Image", + "RelightSimple": "Relight (Simple)", "RemapRange": "Remap Range", "ShuffleChannels": "Shuffle Channels", "Tonemap": "Tonemap",