Added simple bg swap node, added automatic calculation of threshold, included shadow layer in final output, added node to convert to LAB color space, Added node to normalize array layers, added code to rescale histograms based on min and max values

This commit is contained in:
aravind
2024-04-08 23:11:03 +05:30
parent c1b47aa0d1
commit dad3bc5bee
2 changed files with 195 additions and 0 deletions
+6
View File
@@ -13,6 +13,7 @@ from PIL import Image, ImageOps
sys.path.append(tri3d_custom_nodes_path)
from scaled_paste import main_scaled_paste
from simple_bg_swap import (simple_bg_swap, get_threshold_for_bg_swap, RGB_2_LAB, LAB_2_RGB, get_mean_and_standard_deviation, renormalize_array)
from distribution_reshape import (simple_rescale_histogram, get_histogram_limits)
def from_torch_image(image):
@@ -2804,9 +2805,12 @@ NODE_CLASS_MAPPINGS = {
'tri3d-LAB_2_RGB': LAB_2_RGB,
'tri3d-get_mean_and_standard_deviation': get_mean_and_standard_deviation,
'tri3d-renormalize_array': renormalize_array,
"tri3d-simple_rescale_histogram": simple_rescale_histogram,
"tri3d-get_histogram_limits": get_histogram_limits,
}
VERSION = "2.9.0"
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -2844,4 +2848,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
'tri3d-LAB_2_RGB': 'Convert LAB color space to RGB color space' + " v" + VERSION,
'tri3d-get_mean_and_standard_deviation': 'Get mean and standard deviation of array' + " v" + VERSION,
'tri3d-renormalize_array': 'Renormalize the layer to have the given mean and standard deviation' + " v" + VERSION,
"tri3d-simple_rescale_histogram": 'Rescale the layer to have given max and min values' + " v" + VERSION,
"tri3d-get_histogram_limits": 'Calculate max and min values for rescaling histogram' + " v" + VERSION,
}
+189
View File
@@ -0,0 +1,189 @@
import cv2
import os
import torch
import numpy as np
def from_torch_image(image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def get_histogram(array):
array = array.flatten().astype(dtype=np.float64)
hist = np.histogram(array, bins=256, range=(0, 256))
array = hist[0].astype(dtype=np.float64)
array /= len(array)
return array
def get_limits(array, threshold_fraction):
array = get_histogram(array)
left_sum = 0
right_sum = 0
left_start = 0
right_start = len(array) - 1
for i in range(len(array)):
left_index = i
right_index = len(array) - i - 1
left_sum += array[left_index]
right_sum += array[right_index]
if left_sum < threshold_fraction:
left_start = left_index
if right_sum < threshold_fraction:
right_start = right_index
if (left_sum > threshold_fraction) and (right_sum
> threshold_fraction):
return (left_start, right_start)
def do_rescale(x, y1, y2, x1, x2):
x = x.astype(dtype=np.float64)
if x1 > x2:
x1, x2 = x2, x1
if y1 > y2:
y1, y2 = y2, y1
epsilon = 0.0001
y = (x - x1)
y /= (x2 - x1 + epsilon)
y *= (y2 - y1)
y += y1
y = np.clip(y, y1, y2)
for iy in range(y.shape[0]):
for ix in range(y.shape[1]):
if y[iy, ix] > 255:
print(iy, ix)
y = y.astype(dtype=np.uint8)
return y
class get_histogram_limits:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"luminosity_as_mask": ("MASK", ),
"threshold_fraction": ("FLOAT", {
"default": 0.001,
"min": 0.0,
"max": 0.5,
"step": 0.00001,
"round": 0.000001,
"display": "number"
}),
},
}
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("histogram lower limit (x1) as INT",
"histogram upper limit (x2) as INT")
FUNCTION = "test"
#OUTPUT_NODE = False
CATEGORY = "TRI3D"
def test(self, luminosity_as_mask, threshold_fraction):
luminosity_as_mask = from_torch_image(image=luminosity_as_mask)
(left_start,
right_start) = get_limits(array=luminosity_as_mask[0],
threshold_fraction=threshold_fraction)
return (left_start, right_start)
class simple_rescale_histogram:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"layer_as_mask": ("MASK", ),
"y1": ("INT", {
"default": 100,
"min": 0,
"max": 255,
"step": 1,
"display": "number"
}),
"y2": ("INT", {
"default": 200,
"min": 0,
"max": 255,
"step": 1,
"display": "number"
}),
"x1": ("INT", {
"default": 100,
"min": 0,
"max": 255,
"step": 1,
"display": "number"
}),
"x2": ("INT", {
"default": 200,
"min": 0,
"max": 255,
"step": 1,
"display": "number"
})
},
}
RETURN_TYPES = ("MASK", )
RETURN_NAMES = ("rescaled layer as MASK", )
FUNCTION = "test"
CATEGORY = "TRI3D"
def test(self, layer_as_mask, y1, y2, x1, x2):
layer_as_mask = from_torch_image(image=layer_as_mask[0])
layer_as_mask = do_rescale(x=layer_as_mask, y1=y1, y2=y2, x1=x1, x2=x2)
layer_as_mask = to_torch_image(image=layer_as_mask)
layer_as_mask = layer_as_mask.unsqueeze(0)
return (layer_as_mask, )
NODE_CLASS_MAPPINGS = {
"get_histogram_limits": get_histogram_limits,
'simple_rescale_histogram': simple_rescale_histogram
}
NODE_DISPLAY_NAME_MAPPINGS = {
"get_histogram_limits": "get_histogram_limits",
"simple_rescale_histogram": "simple_rescale_histogram"
}