added shuffle, colormatch

This commit is contained in:
spacepxl
2024-02-21 01:50:55 -05:00
parent 2c85cacec0
commit 4fbb0b85ef
+101
View File
@@ -255,6 +255,60 @@ class BlurMaskFast:
return (torch.from_numpy(dup),)
def cv_blur_tensor(images, dx, dy):
if min(dx, dy) > 100:
np_img = torch.nn.functional.interpolate(images.detach().clone().movedim(-1,1), scale_factor=0.1, mode='bilinear').movedim(1,-1).cpu().numpy()
for index, image in enumerate(np_img):
np_img[index] = cv2.GaussianBlur(image, (dx // 20 * 2 + 1, dy // 20 * 2 + 1), 0)
return torch.nn.functional.interpolate(torch.from_numpy(np_img).movedim(-1,1), size=(images.shape[1], images.shape[2]), mode='bilinear').movedim(1,-1)
else:
np_img = images.detach().clone().cpu().numpy()
for index, image in enumerate(np_img):
np_img[index] = cv2.GaussianBlur(image, (dx, dy), 0)
return torch.from_numpy(np_img)
class ColorMatch:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE", ),
"reference": ("IMAGE", ),
"blur": ("INT", {"default": 0, "min": 0, "max": 1023}),
"factor": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01, "round": 0.01}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "batch_normalize"
CATEGORY = "image/filters"
def batch_normalize(self, images, reference, blur, factor):
t = images.detach().clone()
ref = reference.detach().clone()
if ref.shape[0] < t.shape[0]:
ref = ref[0].unsqueeze(0).repeat(t.shape[0], 1, 1, 1)
if blur == 0:
mean = torch.mean(t, (1,2), keepdim=True)
mean_ref = torch.mean(ref, (1,2), keepdim=True)
for i in range(t.shape[0]):
for c in range(3):
t[i,:,:,c] /= mean[i,0,0,c]
t[i,:,:,c] *= mean_ref[i,0,0,c]
else:
d = blur * 2 + 1
blurred = cv_blur_tensor(torch.clamp(t, 0.001, 1), d, d)
blurred_ref = cv_blur_tensor(torch.clamp(ref, 0.001, 1), d, d)
for i in range(t.shape[0]):
for c in range(3):
t[i,:,:,c] /= blurred[i,:,:,c]
t[i,:,:,c] *= blurred_ref[i,:,:,c]
t = torch.lerp(images, t, factor)
return (t,)
class DilateErodeMask:
def __init__(self):
pass
@@ -449,6 +503,49 @@ class RemapRange:
return (torch.from_numpy(i_dup),)
Channel_List = ["red", "green", "blue", "alpha", "white", "black"]
Alpha_List = ["red", "green", "blue", "alpha", "white", "black", "none"]
class ShuffleChannels:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"red": (Channel_List, {"default": "red"}),
"green": (Channel_List, {"default": "green"}),
"blue": (Channel_List, {"default": "blue"}),
"alpha": (Alpha_List, {"default": "none"}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "shuffle"
CATEGORY = "image/filters"
def shuffle(self, image, red, green, blue, alpha):
ch = 3 if alpha == "none" else 4
t = torch.zeros((image.shape[0], image.shape[1], image.shape[2], ch), dtype=image.dtype, device=image.device)
image_copy = image.detach().clone()
ch_key = [red, green, blue, alpha]
for i in range(ch):
if ch_key[i] == "white":
t[:,:,:,i] = 1
elif ch_key[i] == "red":
t[:,:,:,i] = image_copy[:,:,:,0]
elif ch_key[i] == "green":
t[:,:,:,i] = image_copy[:,:,:,1]
elif ch_key[i] == "blue":
t[:,:,:,i] = image_copy[:,:,:,2]
elif ch_key[i] == "alpha":
if image.shape[3] > 3:
t[:,:,:,i] = image_copy[:,:,:,3]
else:
t[:,:,:,i] = 1
return(t,)
class ClampOutliers:
def __init__(self):
pass
@@ -1248,6 +1345,7 @@ NODE_CLASS_MAPPINGS = {
"BlurImageFast": BlurImageFast,
"BlurMaskFast": BlurMaskFast,
"ClampOutliers": ClampOutliers,
"ColorMatch": ColorMatch,
"ConvertNormals": ConvertNormals,
"DifferenceChecker": DifferenceChecker,
"DilateErodeMask": DilateErodeMask,
@@ -1262,6 +1360,7 @@ NODE_CLASS_MAPPINGS = {
"NormalMapSimple": NormalMapSimple,
"OffsetLatentImage": OffsetLatentImage,
"RemapRange": RemapRange,
"ShuffleChannels": ShuffleChannels,
"Tonemap": Tonemap,
"UnJitterImage": UnJitterImage,
"UnTonemap": UnTonemap,
@@ -1280,6 +1379,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"BlurImageFast": "Blur Image (Fast)",
"BlurMaskFast": "Blur Mask (Fast)",
"ClampOutliers": "Clamp Outliers",
"ColorMatch": "Color Match",
"ConvertNormals": "Convert Normals",
"DifferenceChecker": "Difference Checker",
"DilateErodeMask": "Dilate/Erode Mask",
@@ -1294,6 +1394,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"NormalMapSimple": "Normal Map (Simple)",
"OffsetLatentImage": "Offset Latent Image",
"RemapRange": "Remap Range",
"ShuffleChannels": "Shuffle Channels",
"Tonemap": "Tonemap",
"UnJitterImage": "Un-Jitter Image",
"UnTonemap": "UnTonemap",