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

This commit is contained in:
aravind
2024-04-04 18:20:17 +05:30
parent 41ddedb0ba
commit c1b47aa0d1
2 changed files with 34 additions and 36 deletions
+8 -1
View File
@@ -12,7 +12,7 @@ import folder_paths
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)
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)
def from_torch_image(image):
@@ -2801,8 +2801,12 @@ NODE_CLASS_MAPPINGS = {
'tri3d-simple_bg_swap': simple_bg_swap,
'tri3d-get_threshold_for_bg_swap': get_threshold_for_bg_swap,
'tri3d-RGB_2_LAB': RGB_2_LAB,
'tri3d-LAB_2_RGB': LAB_2_RGB,
'tri3d-get_mean_and_standard_deviation': get_mean_and_standard_deviation,
'tri3d-renormalize_array': renormalize_array,
}
VERSION = "2.9.0"
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -2837,4 +2841,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
'tri3d-simple_bg_swap': 'Simple bg swap' + " v" + VERSION,
'tri3d-get_threshold_for_bg_swap': 'Get threshold for bg swap' + " v" + VERSION,
'tri3d-RGB_2_LAB': 'Convert to LAB color space' + " v" + VERSION,
'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,
}
+26 -35
View File
@@ -177,7 +177,6 @@ def get_mu_sigma(array_input, mask_input):
array_input = array_input.astype(dtype=np.float32).flatten()
mask_input = mask_input.astype(dtype=np.float32).flatten()
mask_input /= 255
sum = np.sum(mask_input)
mean = np.sum(array_input * mask_input) / sum
@@ -189,6 +188,17 @@ def get_mu_sigma(array_input, mask_input):
return mean, sigma
def renormalize_array_main(array_input, mask_input, mu, sigma):
array_input_original = array_input.copy()
in_mu, in_sigma = get_mu_sigma(array_input, mask_input)
array_input = (((array_input - in_mu) / in_sigma) * sigma) + mu
array_input_original = (array_input_original *
(1 - mask_input)) + (array_input * mask_input)
return array_input_original
#!/usr/bin/python3
class simple_bg_swap:
@@ -507,15 +517,16 @@ class get_mean_and_standard_deviation:
def test(self, input_array, input_mask):
input_array = from_torch_image(image=input_array)
input_mask = from_torch_image(image=input_mask)
input_array = input_array.cpu().numpy()
input_mask = input_mask.cpu().numpy()
mean, sigma = get_mu_sigma(array_input=input_array[0],
mask_input=input_mask[0])
print(mean, sigma)
return (mean, sigma)
return (
mean,
sigma,
)
class renormalize_array:
@@ -575,38 +586,20 @@ class renormalize_array:
ret = []
if (input_mask.shape[0] == batch_size) and (input_B.shape[0]
== batch_size):
if input_mask.shape[0] == batch_size:
for i in range(batch_size):
input_L_NP = from_torch_image(image=input_L[i])
input_A_NP = from_torch_image(image=input_A[i])
input_B_NP = from_torch_image(image=input_B[i])
tmp = renormalize_array_main(
array_input=input_array[i].cpu().numpy(),
mask_input=input_mask[i].cpu().numpy(),
mu=input_mean,
sigma=input_standard_deviation)
Y_MAX = input_L_NP.shape[0]
X_MAX = input_L_NP.shape[1]
tmp = torch.from_numpy(tmp)
tmp = tmp.unsqueeze(0)
if (input_A_NP.shape[0]
== Y_MAX) and (input_B_NP.shape[0] == Y_MAX) and (
(input_A_NP.shape[1] == X_MAX) and
(input_B_NP.shape[1] == X_MAX)):
image = np.zeros((Y_MAX, X_MAX, 3), dtype=np.uint8)
image[:, :, 0] = input_L_NP
image[:, :, 1] = input_A_NP
image[:, :, 2] = input_B_NP
image = cv2.cvtColor(image, cv2.COLOR_LAB2RGB)
image = to_torch_image(image).unsqueeze(0)
print('image.shape')
print(image.shape)
ret.append(image)
else:
print('Resolution of different layers donot match')
ret.append(tmp)
else:
@@ -614,6 +607,4 @@ class renormalize_array:
ret = torch.cat(ret, dim=0)
print('ret.shape', ret.shape)
return (ret, )