From c1b47aa0d118aa2ca1220d44607a331bc4c08450 Mon Sep 17 00:00:00 2001 From: aravind Date: Thu, 4 Apr 2024 18:20:17 +0530 Subject: [PATCH] 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 --- __init__.py | 9 ++++++- simple_bg_swap.py | 61 ++++++++++++++++++++--------------------------- 2 files changed, 34 insertions(+), 36 deletions(-) diff --git a/__init__.py b/__init__.py index 1370e7e..ddd65bc 100644 --- a/__init__.py +++ b/__init__.py @@ -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, } diff --git a/simple_bg_swap.py b/simple_bg_swap.py index a9ccede..14fa40a 100644 --- a/simple_bg_swap.py +++ b/simple_bg_swap.py @@ -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, )