Merge pull request #34 from TRI3D-LC/extract-pascal-segmentaion

added extract and position nodes for pascal segmentation
This commit is contained in:
NitishTRI3D
2024-06-24 17:06:02 +05:30
committed by GitHub
+359
View File
@@ -835,6 +835,205 @@ class TRI3DExtractPartsBatch:
print(batch_results.shape, "batch_results.shape")
return (batch_results, batch_secondaries)
class TRI3DExtractPascalPartsBatch:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"batch_images": ("IMAGE", ),
"batch_segs": ("IMAGE", ),
"batch_secondaries": ("IMAGE", ),
"margin": ("INT", {
"default": 15,
"min": 0
}),
"background": ("BOOLEAN", {
"default": False
}),
"head": ("BOOLEAN", {
"default": False
}),
"torso": ("BOOLEAN", {
"default": False
}),
"upper_arms": ("BOOLEAN", {
"default": False
}),
"lower_arms": ("BOOLEAN", {
"default": False
}),
"upper_legs": ("BOOLEAN", {
"default": False
}),
"lower_legs": ("BOOLEAN", {
"default": False
}),
},
}
RETURN_TYPES = (
"IMAGE",
"IMAGE",
)
FUNCTION = "main"
CATEGORY = "TRI3D"
def main(self, batch_images, batch_segs, batch_secondaries, margin,
background, head, torso, upper_arms, lower_arms, upper_legs, lower_legs):
import cv2
import numpy as np
import torch
from pprint import pprint
def get_segment_counts(segm):
# Load the segmentation image
# Reshape the image array to be 2D
reshaped = segm.reshape(-1, segm.shape[-1])
# Find unique vectors and their counts
unique_vectors, counts = np.unique(reshaped,
axis=0,
return_counts=True)
segment_counts = list(zip(unique_vectors, counts))
pprint(segment_counts)
return segment_counts
def bounded_image(seg_img, color_code_list, input_img):
import cv2
import numpy as np
# Create a mask for hands
seg_img = cv2.resize(seg_img,
(input_img.shape[1], input_img.shape[0]),
interpolation=cv2.INTER_NEAREST)
hand_mask = np.zeros_like(seg_img[:, :, 0])
for color in color_code_list:
lowerb = np.array(color, dtype=np.uint8)
upperb = np.array(color, dtype=np.uint8)
temp_mask = cv2.inRange(seg_img, lowerb, upperb)
hand_mask = cv2.bitwise_or(hand_mask, temp_mask)
# Find contours to get the bounding box of the hands
contours, _ = cv2.findContours(hand_mask, cv2.RETR_EXTERNAL,
cv2.CHAIN_APPROX_SIMPLE)
# If no contours were found, just return None
if not contours:
return None
# Combine all contours to find encompassing bounding box
all_points = np.concatenate(contours, axis=0)
x, y, w, h = cv2.boundingRect(all_points)
print(x, y, w, h, "x,y,w,h")
x = max(x - margin, 0)
y = max(y - margin, 0)
# Ensure width does not exceed image boundary
w = min(w + 2 * margin, input_img.shape[1] - x)
# Ensure height does not exceed image boundary
h = min(h + 2 * margin, input_img.shape[0] - y)
print(x, y, w, h, "x,y,w,h")
print(input_img.shape, "input_img.shape")
# Extract the region from the original image that contains both hands
hand_region = input_img[y:y + h, x:x + w]
return hand_region
def tensor_to_cv2_img(tensor, remove_alpha=False):
# This will give us (H, W, C)
i = 255. * tensor.squeeze(0).cpu().numpy()
img = np.clip(i, 0, 255).astype(np.uint8)
return img
def cv2_img_to_tensor(img):
img = img.astype(np.float32) / 255.0
img = torch.from_numpy(img)[
None,
]
return img
batch_results = []
images = []
secondaries = []
# cv2_secondary = tensor_to_cv2_img(batch_secondaries)
for i in range(batch_images.shape[0]):
image = batch_images[i]
seg = batch_segs[i]
cv2_image = tensor_to_cv2_img(image)
cv2_secondary = tensor_to_cv2_img(batch_secondaries[i])
cv2_seg = tensor_to_cv2_img(seg)
color_code_list = []
################# ATR MAPPING#################
if background:
color_code_list.append([0, 0, 0])
if head:
color_code_list.append([128, 0, 0])
if torso:
color_code_list.append([0, 128, 0])
if upper_arms:
color_code_list.append([128, 128, 0])
if lower_arms:
color_code_list.append([0, 0, 128])
if upper_legs:
color_code_list.append([128, 0, 128])
if lower_legs:
color_code_list.append([0, 128, 128])
bimage = bounded_image(cv2_seg, color_code_list, cv2_image)
bsecondary = bounded_image(cv2_seg, color_code_list, cv2_secondary)
# Handle case when bimage is None to avoid error during conversion to tensor
if bimage is not None:
images.append(bimage)
else:
num_channels = cv2_image.shape[2] if len(
cv2_image.shape) > 2 else 1
black_img = np.zeros((10, 10, num_channels),
dtype=cv2_image.dtype)
images.append(black_img)
if bsecondary is not None:
secondaries.append(bsecondary)
else:
num_channels = cv2_image.shape[2] if len(
cv2_image.shape) > 2 else 1
black_img = np.zeros((10, 10, num_channels),
dtype=cv2_image.dtype)
secondaries.append(black_img)
# Get max height and width
max_height = max(img.shape[0] for img in images)
max_width = max(img.shape[1] for img in images)
batch_results = []
batch_secondaries = []
for img in images:
# Resize the image to max height and width
resized_img = cv2.resize(img, (max_width, max_height),
interpolation=cv2.INTER_CUBIC)
tensor_img = cv2_img_to_tensor(resized_img)
batch_results.append(tensor_img.squeeze(0))
for sec in secondaries:
# Resize the image to max height and width
resized_sec = cv2.resize(sec, (max_width, max_height),
interpolation=cv2.INTER_NEAREST)
tensor_sec = cv2_img_to_tensor(resized_sec)
batch_secondaries.append(tensor_sec.squeeze(0))
batch_results = torch.stack(batch_results)
batch_secondaries = torch.stack(batch_secondaries)
print(batch_results.shape, "batch_results.shape")
return (batch_results, batch_secondaries)
class TRI3DPositionPartsBatch:
@@ -1050,6 +1249,162 @@ class TRI3DPositionPartsBatch:
return (batch_results, )
class TRI3DPositionPascalPartsBatch:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"batch_images": ("IMAGE", ),
"batch_segs": ("IMAGE", ),
"batch_handimgs": ("IMAGE", ),
"margin": ("INT", {
"default": 15,
"min": 0
}),
"background": ("BOOLEAN", {
"default": False
}),
"head": ("BOOLEAN", {
"default": False
}),
"torso": ("BOOLEAN", {
"default": False
}),
"upper_arms": ("BOOLEAN", {
"default": False
}),
"lower_arms": ("BOOLEAN", {
"default": False
}),
"upper_legs": ("BOOLEAN", {
"default": False
}),
"lower_legs": ("BOOLEAN", {
"default": False
}),
},
}
RETURN_TYPES = ("IMAGE", )
FUNCTION = "main"
CATEGORY = "TRI3D"
def main(self, batch_images, batch_segs, batch_handimgs, margin,
background, head, torso, upper_arms, lower_arms, upper_legs, lower_legs):
import cv2
import numpy as np
import torch
from pprint import pprint
def bounded_image_points(seg_img, color_code_list, input_img):
import cv2
import numpy as np
# Create a mask for hands
seg_img = cv2.resize(seg_img,
(input_img.shape[1], input_img.shape[0]),
interpolation=cv2.INTER_NEAREST)
hand_mask = np.zeros_like(seg_img[:, :, 0])
for color in color_code_list:
lowerb = np.array(color, dtype=np.uint8)
upperb = np.array(color, dtype=np.uint8)
temp_mask = cv2.inRange(seg_img, lowerb, upperb)
hand_mask = cv2.bitwise_or(hand_mask, temp_mask)
# Find contours to get the bounding box of the hands
contours, _ = cv2.findContours(hand_mask, cv2.RETR_EXTERNAL,
cv2.CHAIN_APPROX_SIMPLE)
# If no contours were found, just return None
if not contours:
return None
# Combine all contours to find encompassing bounding box
all_points = np.concatenate(contours, axis=0)
x, y, w, h = cv2.boundingRect(all_points)
x = max(x - margin, 0)
y = max(y - margin, 0)
# Ensure width does not exceed image boundary
w = min(w + 2 * margin, input_img.shape[1] - x)
# Ensure height does not exceed image boundary
h = min(h + 2 * margin, input_img.shape[0] - y)
return (x, y, w, h)
def tensor_to_cv2_img(tensor, remove_alpha=False):
# This will give us (H, W, C)
i = 255. * tensor.squeeze(0).cpu().numpy()
img = np.clip(i, 0, 255).astype(np.uint8)
return img
def cv2_img_to_tensor(img):
img = img.astype(np.float32) / 255.0
img = torch.from_numpy(img)[
None,
]
return img
def unsharp_mask(image, sigma=1.0, strength=1.0):
# Blur the image
blurred_image = cv2.GaussianBlur(image, (0, 0), sigma)
# Calculate the sharpened image
sharpened_image = cv2.addWeighted(image, 1.0 + strength, blurred_image, -strength, 0)
return sharpened_image
batch_results = []
for i in range(batch_images.shape[0]):
image = batch_images[i]
seg = batch_segs[i]
handimg = batch_handimgs[i]
cv2_image = tensor_to_cv2_img(image)
cv2_seg = tensor_to_cv2_img(seg)
color_code_list = []
################# ATR MAPPING#################
if background:
color_code_list.append([0, 0, 0])
if head:
color_code_list.append([128, 0, 0])
if torso:
color_code_list.append([0, 128, 0])
if upper_arms:
color_code_list.append([128, 128, 0])
if lower_arms:
color_code_list.append([0, 0, 128])
if upper_legs:
color_code_list.append([128, 0, 128])
if lower_legs:
color_code_list.append([0, 128, 128])
positions = bounded_image_points(cv2_seg, color_code_list,
cv2_image)
try:
cv2_handimg = tensor_to_cv2_img(handimg)
cv2_handimg = cv2.resize(cv2_handimg,
(positions[2], positions[3]),
interpolation=cv2.INTER_AREA)
cv2_handimg = unsharp_mask(cv2_handimg)
cv2_image[positions[1]:positions[1] + positions[3],
positions[0]:positions[0] +
positions[2]] = cv2_handimg
except Exception as e:
print(e)
pass
b_tensor_img = cv2_img_to_tensor(cv2_image)
batch_results.append(b_tensor_img.squeeze(0))
batch_results = torch.stack(batch_results)
return (batch_results, )
class TRI3DSwapPixels:
@@ -3157,6 +3512,8 @@ NODE_CLASS_MAPPINGS = {
'tri3d-extract-parts-batch': TRI3DExtractPartsBatch,
"tri3d-extract-parts-batch2": TRI3DExtractPartsBatch2,
"tri3d-position-parts-batch": TRI3DPositionPartsBatch,
'tri3d-extract-pascal-parts-batch': TRI3DExtractPascalPartsBatch,
"tri3d-position-pascal-parts-batch": TRI3DPositionPascalPartsBatch,
"tri3d-swap-pixels": TRI3DSwapPixels,
"tri3d-skin-feathered-padded-mask": TRI3DSkinFeatheredPaddedMask,
"tri3d-interaction-canny": TRI3DInteractionCanny,
@@ -3205,6 +3562,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
'tri3d-extract-parts-batch': 'Extract Parts Batch' + " v" + VERSION,
'tri3d-extract-parts-batch2': 'Extract Parts Batch 2' + " v" + VERSION,
"tri3d-position-parts-batch": "Position Parts Batch" + " v" + VERSION,
'tri3d-extract-pascal-parts-batch': 'Extract Pascal Parts Batch' + " v" + VERSION,
"tri3d-position-pascal-parts-batch": "Position Pascal Parts Batch" + " v" + VERSION,
"tri3d-swap-pixels": "Swap Pixels by Mask" + " v" + VERSION,
"tri3d-skin-feathered-padded-mask":
"Skin Feathered Padded Mask" + " v" + VERSION,