From dad3bc5beeaf75c3df95a959a595d8a6f9ef162f Mon Sep 17 00:00:00 2001 From: aravind Date: Mon, 8 Apr 2024 23:11:03 +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, added code to rescale histograms based on min and max values --- __init__.py | 6 ++ distribution_reshape.py | 189 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 195 insertions(+) create mode 100644 distribution_reshape.py diff --git a/__init__.py b/__init__.py index ddd65bc..ad3a71c 100644 --- a/__init__.py +++ b/__init__.py @@ -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, } diff --git a/distribution_reshape.py b/distribution_reshape.py new file mode 100644 index 0000000..4cf1994 --- /dev/null +++ b/distribution_reshape.py @@ -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" +}