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:
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
Reference in New Issue
Block a user