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" +}