RelightSimple

This commit is contained in:
spacepxl
2024-02-23 04:17:49 -05:00
parent 09c410baec
commit 8517400213
+36
View File
@@ -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",