64 Commits
Author SHA1 Message Date
Ubuntu a57e48c814 levind_abhi 2025-05-23 10:17:25 +00:00
NitishTRI3D bec8344e93 Merge pull request #57 from TRI3D-LC/neck_modify
Neck modify
2025-03-10 11:28:11 +05:30
Ubuntu cf47490051 neck_modify 4.9.0v 2025-03-10 05:57:46 +00:00
Ubuntu 22d65c983c polygon neck 2025-03-07 08:28:35 +00:00
Ubuntu aa7109592e working code for weighted point to ears logic 2025-03-07 06:50:51 +00:00
NitishTRI3D 6a34903aba Merge pull request #56 from TRI3D-LC/narrowfy
added margin as an input to the narrowfy node
2025-02-21 16:37:26 +05:30
Ubuntu 9e8d4a6148 added margin as an input to the narrowfy node 2025-02-21 11:02:46 +00:00
NitishTRI3D 942b41383b Merge pull request #55 from TRI3D-LC/smart_depth
save changes
2025-02-20 11:24:30 +05:30
Ubuntu 5a5fb0d129 save changes 2025-02-20 05:52:52 +00:00
NitishTRI3D 9e9c958862 Merge pull request #54 from TRI3D-LC/smart_depth
Smart depth
2025-02-20 10:58:08 +05:30
Ubuntu cbc31d761e narrowfy 2025-02-17 13:18:44 +00:00
Ubuntu d02fcf8118 narrowfy 2025-02-17 10:39:16 +00:00
NitishTRI3D 8f5cd058fe Merge pull request #53 from TRI3D-LC/image_entend
Image entend
2025-01-29 13:11:00 +05:30
Ubuntu bd3dbad41c modified code to extend image to preserve original aspect ratio 2025-01-27 12:07:43 +00:00
Ubuntu 15ae5d9fef rectified bug of not importing ratio and commented print statements 2025-01-27 12:02:54 +00:00
Ubuntu b430d01c5b taking ratio as an input 2025-01-27 11:48:53 +00:00
Ubuntu d10c1195e6 TRI3D_Image_extend node to extend image for a close up image input 2025-01-27 10:57:51 +00:00
Ubuntu c68655e6c3 v4.8.5; tri3d nsfw 2025-01-24 13:35:08 +00:00
Ubuntu 110585389e v4.8.5; tri3d nsfw 2025-01-24 12:10:20 +00:00
NitishTRI3D 23ac5cb1c7 Merge pull request #52 from TRI3D-LC/smartbox_neck
smartbox_neck, 4.8.4
2025-01-24 08:50:59 +05:30
Ubuntu 96b1198824 smartbox_neck, 4.8.4 2025-01-22 07:19:47 +00:00
Ubuntu c1a24a2244 v4.8.3 ; hip calculation skipping negative 2025-01-17 14:18:06 +00:00
Ubuntu 77b4f2713f v4.8.3 ; hip calculation skipping negative 2025-01-17 14:01:22 +00:00
Ubuntu 8b00fcffec v4.8.2.1 , highest of lower hip points 2025-01-17 08:59:07 +00:00
Ubuntu edb141157f skip head mask, 4.8.2 2025-01-17 03:21:26 +00:00
Ubuntu a219efcd14 Merge branch 'main' of https://github.com/TRI3D-LC/tri3d-comfyui-nodes 2025-01-16 09:32:11 +00:00
Ubuntu 62e0b4b8ba skip head 2025-01-16 09:31:56 +00:00
Ubuntu fc55569c56 hip not found error 2025-01-10 08:50:01 +00:00
NitishTRI3D b19de8ac31 Merge pull request #51 from TRI3D-LC/check
correct import
2025-01-08 19:01:56 +05:30
Ubuntu f19063af5b correct import 2025-01-08 13:30:05 +00:00
NitishTRI3D 358118e6d2 Merge pull request #50 from TRI3D-LC/check
correct import
2025-01-08 18:36:49 +05:30
Ubuntu 0d6a998eca correct import 2025-01-08 13:05:43 +00:00
Ubuntu 2cdd0c43bf importing 2025-01-08 12:45:38 +00:00
NitishTRI3D 1a2d9de309 Merge pull request #49 from TRI3D-LC/ahead
Ahead
2025-01-08 18:09:53 +05:30
Ubuntu c0b0e48b1c Merge branch 'main' of https://github.com/TRI3D-LC/tri3d-comfyui-nodes 2025-01-08 12:35:51 +00:00
Ubuntu 9d2b368bb6 4.8 smart box release 2025-01-08 12:34:17 +00:00
Ubuntu a4f3b113e1 Merge branch 'main' of https://github.com/TRI3D-LC/tri3d-comfyui-nodes 2024-09-26 05:27:28 +00:00
NitishTRI3D 07eb4d19ed Merge pull request #48 from TRI3D-LC/image_stack
Added image stacking node
2024-09-26 10:57:16 +05:30
aravindhv10 cbfbb79ad9 Added qwen stuff 2024-09-25 15:09:57 +05:30
aravindhv10 95bccde3ad Added image stacking node 2024-09-18 13:15:46 +05:30
aravindhv10 b8f7f78466 Added image stacking node 2024-09-18 12:16:13 +05:30
aravindhv10 833473e39f Added image stacking node 2024-09-17 14:03:19 +05:30
Ubuntu 1a44ee657e saving merge 2024-09-05 18:52:32 +00:00
NitishTRI3D 52f4ad7854 Merge pull request #47 from TRI3D-LC/facer-mask-extractor
added new node to extract mask from facer
2024-08-26 12:53:17 +05:30
Ubuntu e31a4346f5 added new node to extract mask from facer 2024-08-24 14:50:29 +00:00
NitishTRI3D 7fd309086d Merge pull request #46 from TRI3D-LC/only_trouser
Only trouser
2024-08-22 19:26:12 +05:30
Ubuntu 5c1a8a6dfa Merge branch 'recolor_changes' into only_trouser 2024-08-22 13:39:03 +00:00
Ubuntu c4ab544fd1 added in tri3d_recolor_lab 2024-08-22 13:38:48 +00:00
Ubuntu ff2cd75e6f added node to detect trouser images 2024-08-22 12:55:16 +00:00
Apple a6cfa6483e v471, original sigma used 2024-08-22 12:35:49 +05:30
Ubuntu 09034883c8 v4.7.1; keeping same standard deviation in recoloring LAB 2024-08-22 06:59:34 +00:00
NitishTRI3D e61adaf421 Merge pull request #45 from TRI3D-LC/stagger_mask_fix
fixed satggered masking in fill mask node
2024-08-12 11:57:09 +05:30
Ubuntu 39f69cc278 fixed satggered masking in fill mask node 2024-08-12 06:21:59 +00:00
NitishTRI3D 40f2d4bc55 Merge pull request #44 from TRI3D-LC/comping-fixes
changes in utility node
2024-08-07 14:58:13 +05:30
Ubuntu fcfb331421 chnaged version 2024-08-07 09:27:04 +00:00
Ubuntu ab939e704b changes in utility node 2024-08-06 10:34:07 +00:00
Apple 6c89ebdc5f v4.5 2024-07-31 16:44:46 +05:30
NitishTRI3D 818033bf64 Merge pull request #43 from TRI3D-LC/utility_node
Utility nodes
2024-07-31 16:43:48 +05:30
Ubuntu b79e8eaa1a utility node changes 2024-07-31 08:29:37 +00:00
Ubuntu c4d1a3276f changed clean mask node 2024-07-30 17:15:12 +00:00
Ubuntu 3fe6e7e0f6 added nodes to extract and position part of image, from pose 2024-07-29 15:30:41 +00:00
Ubuntu ec79d8057f some changes 2024-07-29 13:03:08 +00:00
Ubuntu 3472ca0ad3 initial commit 2024-07-29 12:42:10 +00:00
Ubuntu 4d681027fb added few utility nodes 2024-07-29 07:27:19 +00:00
11 changed files with 1937 additions and 37 deletions
+2
View File
@@ -11,3 +11,5 @@ cloth-segmentation/model/cloth_segm.pth
dwpose/keypoints/
huggingface/
safetychecker/model.safetensors
+87 -20
View File
@@ -15,12 +15,22 @@ from scaled_paste import main_scaled_paste
from scaled_paste import main_scaled_paste_2
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)
from utility_nodes import TRI3D_clean_mask, TRI3D_extract_pose_part, TRI3D_position_pose_part, TRI3D_fill_mask, TRI3D_is_only_trouser, TRI3D_extract_facer_mask
from utility_nodes import TRI3D_extract_facer_mask
from .AEMatter import (load_AEMatter_Model, run_AEMatter_inference)
from .light_layer import main_light_layer
from .image_stack import (
H_Stack_Images,
SaveImage_absolute,
SaveText_absolute,
Wait_And_Read_File,
)
def from_torch_image(image):
image = image.squeeze().cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
@@ -228,7 +238,7 @@ class TRI3DLEVINDABHICLOTHSEGBATCH:
},
}
RETURN_TYPES = ("IMAGE", )
RETURN_TYPES = ("IMAGE", "IMAGE", "IMAGE")
FUNCTION = "main"
CATEGORY = "TRI3D"
@@ -282,21 +292,33 @@ class TRI3DLEVINDABHICLOTHSEGBATCH:
# Collect and return the results
batch_results = []
mask0_batch = []
mask1_batch = []
mask2_batch = []
for i in range(images.shape[0]):
cv2_segm = cv2.imread(LSEG_OUTPUT_PATH + f'image{i}.png', cv2.IMREAD_UNCHANGED) # Read PNG with alpha channel
cv2_segm = cv2.cvtColor(cv2_segm, cv2.COLOR_BGRA2RGBA) # Convert from BGRA to RGBA
b_tensor_img = cv2_img_to_tensor(cv2_segm)
batch_results.append(b_tensor_img.squeeze(0))
batch_results = torch.stack(batch_results)
return (batch_results, )
mask0_path = os.path.join(LSEG_OUTPUT_PATH, f"{i}__mask0.png")
mask1_path = os.path.join(LSEG_OUTPUT_PATH, f"{i}__mask1.png")
mask2_path = os.path.join(LSEG_OUTPUT_PATH, f"{i}__mask2.png")
mask0_img = cv2.imread(mask0_path, cv2.IMREAD_UNCHANGED)
mask1_img = cv2.imread(mask1_path, cv2.IMREAD_UNCHANGED)
mask2_img = cv2.imread(mask2_path, cv2.IMREAD_UNCHANGED)
# Ensure single channel, convert to 3 channel if needed for consistency
if mask0_img is not None and len(mask0_img.shape) == 2:
mask0_img = cv2.cvtColor(mask0_img, cv2.COLOR_GRAY2RGB)
if mask1_img is not None and len(mask1_img.shape) == 2:
mask1_img = cv2.cvtColor(mask1_img, cv2.COLOR_GRAY2RGB)
if mask2_img is not None and len(mask2_img.shape) == 2:
mask2_img = cv2.cvtColor(mask2_img, cv2.COLOR_GRAY2RGB)
mask0_tensor = cv2_img_to_tensor(mask0_img).squeeze(0)
mask1_tensor = cv2_img_to_tensor(mask1_img).squeeze(0)
mask2_tensor = cv2_img_to_tensor(mask2_img).squeeze(0)
mask0_batch.append(mask0_tensor)
mask1_batch.append(mask1_tensor)
mask2_batch.append(mask2_tensor)
mask0_batch = torch.stack(mask0_batch)
mask1_batch = torch.stack(mask1_batch)
mask2_batch = torch.stack(mask2_batch)
return (mask0_batch, mask1_batch, mask2_batch)
@@ -1931,6 +1953,8 @@ class TRI3DDWPose_Preprocessor:
out_image = torch.stack(out_image_list, dim=0)
del model
# print(save_file_path, "save_file_path")
return (out_image, save_file_path)
@@ -2656,9 +2680,12 @@ class TRI3D_reLUM:
mu_2, sigma_2 = get_mu_sigma(array_input=image_2[:, :, i],
mask_input=mask_2)
sigma_calculated = sigma_1 * factor_sigma[i]
if factor_sigma[i] < 0:
sigma_calculated = sigma_2
image_2[:, :, i] = (
((image_2[:, :, i] - mu_2) / sigma_2) *
(sigma_1 * factor_sigma[i])) + (mu_1 * factor_mean[i])
sigma_calculated) + (mu_1 * factor_mean[i])
image_2 = np.clip(image_2, 0, 255)
image_2 = image_2.astype(dtype=np.uint8)
@@ -2833,9 +2860,14 @@ class TRI3D_recolor_LAB:
mu_2, sigma_2 = get_mu_sigma(array_input=image_2[:, :, i],
mask_input=mask_2)
sigma_calculated = sigma_1 * factor_sigma[i]
if factor_sigma[i] < 0:
sigma_calculated = sigma_2
image_2[:, :, i] = (
((image_2[:, :, i] - mu_2) / sigma_2) *
(sigma_1 * factor_sigma[i])) + (mu_1 * factor_mean[i])
sigma_calculated) + (mu_1 * factor_mean[i])
image_2 = np.clip(image_2, 0, 255)
image_2 = image_2.astype(dtype=np.uint8)
@@ -3666,7 +3698,6 @@ class TRI3D_BGREMOVE_MEGA():
batch_results = torch.stack(batch_results)
batch_results_masks = torch.stack(batch_results_masks)
return (batch_results,batch_results_masks)
@@ -3675,6 +3706,8 @@ class TRI3D_BGREMOVE_MEGA():
from photoroom import TRI3D_photoroom_bgremove_api
from smart_box import TRI3D_SmartBox, TRI3D_Skip_HeadMask, TRI3D_Skip_HeadMask_AddNeck, TRI3D_Image_extend, TRI3D_Smart_Depth, TRI3D_NarrowfyImage
from nsfw import TRI3DNSFWFilter
# A dictionary that contains all nodes you want to export with their names
# NOTE: names should be globally unique
@@ -3725,10 +3758,27 @@ NODE_CLASS_MAPPINGS = {
'tri3d-run_AEMatter_inference': run_AEMatter_inference,
"tri3d-bgremove-mega" :TRI3D_BGREMOVE_MEGA,
'tri3d-flexible_color_extract' : main_light_layer,
'tri3d-clean_mask': TRI3D_clean_mask,
"tri3d-extract_pose_part": TRI3D_extract_pose_part,
"tri3d_position_pose_part":TRI3D_position_pose_part,
"tri3d_fill_mask": TRI3D_fill_mask,
"tri3d_is_only_trouser": TRI3D_is_only_trouser,
"tri3d_extract_facer_mask":TRI3D_extract_facer_mask,
"tri3d_H_Stack_Images": H_Stack_Images,
"tri3d_SaveImage_absolute":SaveImage_absolute,
"tri3d_SaveText_absolute":SaveText_absolute,
"tri3d_Wait_And_Read_File":Wait_And_Read_File,
"tri3d_SmartBox": TRI3D_SmartBox,
"tri3d_Skip_HeadMask": TRI3D_Skip_HeadMask,
"tri3d_Skip_HeadMask_AddNeck": TRI3D_Skip_HeadMask_AddNeck,
"tri3d_Image_extend": TRI3D_Image_extend,
"tri3d_Smart_Depth": TRI3D_Smart_Depth,
"tri3d_NSFWFilter": TRI3DNSFWFilter,
"tri3d_NarrowfyImage": TRI3D_NarrowfyImage,
}
VERSION = "4.4"
VERSION = "4.9.0"
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
"tri3d-photoroom-bgremove-api": "Photoroom BG Remove" + " v" + VERSION,
@@ -3779,4 +3829,21 @@ NODE_DISPLAY_NAME_MAPPINGS = {
'tri3d-run_AEMatter_inference': 'Run AEMatter inference' + ' v' + VERSION,
"tri3d-bgremove-mega": "BG Remove Mega" + " v" + VERSION,
'tri3d-flexible_color_extract': "Flexible color extract" + " v" + VERSION,
'tri3d-clean_mask': "Clear small patches" + " v" + VERSION,
"tri3d-extract_pose_part": "Extract pose part" + " v" + VERSION,
"tri3d_position_pose_part": "Position pose part" + " v" + VERSION,
"tri3d_fill_mask": "Fill mask" + " v" + VERSION,
"tri3d_is_only_trouser": "Is only trouser" + " v" + VERSION,
"tri3d_extract_facer_mask": "Extract facer mask" + " v" + VERSION,
"tri3d_H_Stack_Images": "Stack images for cat vton with flux" + " v" + VERSION,
"tri3d_SaveImage_absolute": "Save image to an absolute path and provide text optional to control execution order" + " v" + VERSION,
"tri3d_SaveText_absolute": "Save text to an absolute path and provide text optional to control execution order " + " v" + VERSION,
"tri3d_Wait_And_Read_File": "Wait and read text file, optional control from text " + " v" + VERSION,
"tri3d_SmartBox": "Smart Box" + " v" + VERSION,
"tri3d_Skip_HeadMask": "Skip Head Mask" + " v" + VERSION,
"tri3d_Skip_HeadMask_AddNeck": "Skip Head Mask and add neck" + " v" + VERSION,
"tri3d_NSFWFilter": "TRI3D NSFW Filter" + " v" + VERSION,
"tri3d_Image_extend": "Image extend" + " v" + VERSION,
"tri3d_Smart_Depth": "Smart Depth" + " v" + VERSION,
"tri3d_NarrowfyImage": "Narrowfy Image" + " v" + VERSION,
}
+24 -6
View File
@@ -14,16 +14,34 @@ def initialize_and_load_models():
net = initialize_and_load_models()
def run(img):
def run(img, image_id, output_dir):
palette = get_palette(4)
cloth_seg = generate_mask(img, net=net,device=device)
return cloth_seg
mask0, mask1, mask2, cloth_seg = generate_mask(img, net=net, device=device, image_id=image_id, output_dir=output_dir)
return mask0, mask1, mask2, cloth_seg
INPUT_PATH = "./input/"
OUTPUT_PATH = "./output/"
import os
for cur_image in os.listdir(INPUT_PATH):
mask0_paths = []
mask1_paths = []
mask2_paths = []
cloth_paths = []
for idx, cur_image in enumerate(os.listdir(INPUT_PATH)):
img = PIL.Image.open(INPUT_PATH + cur_image)
cloth_seg = run(img)
cloth_seg.save(OUTPUT_PATH + cur_image, format="PNG")
mask0, mask1, mask2, cloth_seg = run(img, image_id=idx, output_dir=OUTPUT_PATH)
# Save masks and cloth_seg with unique names (already saved in generate_mask)
mask0_path = os.path.join(OUTPUT_PATH, f"{idx}__mask0.png")
mask1_path = os.path.join(OUTPUT_PATH, f"{idx}__mask1.png")
mask2_path = os.path.join(OUTPUT_PATH, f"{idx}__mask2.png")
cloth_path = os.path.join(OUTPUT_PATH, f"{idx}__extracted_garment.png")
mask0_paths.append(mask0_path)
mask1_paths.append(mask1_path)
mask2_paths.append(mask2_path)
cloth_paths.append(cloth_path)
print("Mask0 batch:", mask0_paths)
print("Mask1 batch:", mask1_paths)
print("Mask2 batch:", mask2_paths)
print("Garment batch:", cloth_paths)
+32 -10
View File
@@ -101,14 +101,16 @@ def apply_transform(img):
from PIL import Image
def generate_mask(input_image, net, device='cpu'):
def generate_mask(input_image, net, device='cpu', image_id=None, output_dir=None):
img = input_image
img_size = img.size
img = img.resize((768, 768), Image.BICUBIC)
image_tensor = apply_transform(img)
image_tensor = torch.unsqueeze(image_tensor, 0)
output_dir = os.path.join(opt.output, 'extracted_garment')
# Allow output_dir override for batch processing
if output_dir is None:
output_dir = os.path.join(opt.output, 'extracted_garment')
os.makedirs(output_dir, exist_ok=True)
with torch.no_grad():
@@ -118,30 +120,50 @@ def generate_mask(input_image, net, device='cpu'):
output_tensor = torch.squeeze(output_tensor, dim=0)
output_arr = output_tensor.cpu().numpy()
# Create and save individual masks for classes 1, 2, 3
classes_of_interest = [1, 2, 3]
mask_imgs = []
for idx, cls in enumerate(classes_of_interest):
mask = np.zeros_like(output_arr, dtype=np.uint8)
mask[output_arr == cls] = 255
if mask.ndim > 2:
mask = mask.squeeze()
if mask.ndim != 2:
raise ValueError(f"mask{idx} must be a 2-dimensional array")
mask_img = Image.fromarray(mask, mode='L').resize(img_size, Image.BICUBIC)
# Save with unique name if image_id is provided
if image_id is not None:
mask_path = os.path.join(output_dir, f'{image_id}__mask{idx}.png')
else:
mask_path = os.path.join(output_dir, f'mask{idx}.png')
mask_img.save(mask_path, format="PNG")
print(f"Saved mask{idx} at: {mask_path}")
mask_imgs.append(mask_img)
# Create a binary mask where selected classes are 1, others are 0
binary_mask = np.zeros_like(output_arr, dtype=np.uint8)
classes_of_interest = [1, 2, 3] # Modify this list according to your classes of interest
for cls in classes_of_interest:
binary_mask[output_arr == cls] = 255
# Ensure binary_mask is 2D
if binary_mask.ndim > 2:
binary_mask = binary_mask.squeeze() # Removes single-dimensional entries from the shape
binary_mask = binary_mask.squeeze()
if binary_mask.ndim != 2:
raise ValueError("binary_mask must be a 2-dimensional array")
binary_mask_img = Image.fromarray(binary_mask, mode='L').resize(img_size, Image.BICUBIC)
# Create an RGBA image for the output
extracted_garment = Image.new("RGBA", img_size)
original_img = img.resize(img_size) # Resize the processed image back to original size
original_img = img.resize(img_size)
extracted_garment.paste(original_img, mask=binary_mask_img)
# Save the garment image with transparency
garment_path = os.path.join(output_dir, 'extracted_garment.png')
if image_id is not None:
garment_path = os.path.join(output_dir, f'{image_id}__extracted_garment.png')
else:
garment_path = os.path.join(output_dir, 'extracted_garment.png')
extracted_garment.save(garment_path, format="PNG")
print(f"Saved extracted garment at: {garment_path}")
return extracted_garment
return (*mask_imgs, extracted_garment)
# def generate_mask(input_image, net, device='cpu'):
# img = input_image
+3 -1
View File
@@ -274,4 +274,6 @@ def switch_to_backpose(input_keypoints, input_width):
x,y = input_keypoints[i]
input_keypoints[i] = [input_width - x, y]
return input_keypoints
return input_keypoints
Executable
+203
View File
@@ -0,0 +1,203 @@
#!/usr/bin/python3
from PIL import Image, ImageOps, ImageSequence, ImageFile
from PIL.PngImagePlugin import PngInfo
import cv2
import hashlib
import json
import logging
import math
import numpy as np
import os
import random
import safetensors.torch
import sys
import time
import torch
import traceback
def load_image(path):
return torch.from_numpy(cv2.imread(
path, cv2.IMREAD_COLOR)).to(dtype=torch.float32) / 255.0
def do_stack(img1, img2):
dim = max(max(img1.shape[0], img2.shape[0]), img1.shape[1] + img2.shape[1])
out = torch.zeros((dim, dim, 3), dtype=img1.dtype, device=img1.device) + 1
diff1 = (out.shape[0] - img1.shape[0]) // 2
diff2 = (out.shape[0] - img2.shape[0]) // 2
part0 = 0
part1 = img1.shape[1]
part2 = img2.shape[1] + img1.shape[1]
out[diff1:diff1 + img1.shape[0], part0:part1, :] = img1
out[diff2:diff2 + img2.shape[0], part1:part2, :] = img2
return out
def save_image(image, outpath):
cv2.imwrite(outpath,
(image * 255).to(dtype=torch.uint8).detach().cpu().numpy())
class H_Stack_Images:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image_L": ("IMAGE", ),
"image_R": ("IMAGE", ),
},
}
RETURN_TYPES = ("IMAGE", )
FUNCTION = "test"
CATEGORY = "TRI3D"
def test(self, image_L, image_R):
return (do_stack(img1=image_L[0], img2=image_R[0]).unsqueeze(0), )
class SaveImage_absolute:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE", {
"tooltip": "The images to save."
}),
"absolute_filename": ("STRING", {
"default":
"image.png",
"tooltip":
"The absolute path to the file to save."
})
},
}
RETURN_TYPES = ("STRING", )
RETURN_NAMES = ("text to control order", )
FUNCTION = "save_images"
OUTPUT_NODE = True
CATEGORY = "image"
DESCRIPTION = "Saves the input images to an absolute path."
def save_images(self, images, absolute_filename):
i = 255.0 * images[0].cpu().numpy()
img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
img.save(absolute_filename)
return (absolute_filename, )
class SaveText_absolute:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {
"multiline": True,
"dynamicPrompts": True,
"tooltip": "Text to be saved to the file."
}),
"absolute_filename": ("STRING", {
"default":
"image.txt",
"tooltip":
"The absolute path to the file to save."
})
},
"optional": {
"text_opt": ("STRING", {
"multiline":
True,
"dynamicPrompts":
True,
"tooltip":
"Text to provide order when necessary (to create work files after txt files)."
}),
}
}
RETURN_TYPES = ("STRING", )
RETURN_NAMES = ("same text as input", )
FUNCTION = "save_text"
OUTPUT_NODE = True
CATEGORY = "text"
DESCRIPTION = "Saves the input text to an absolute path."
def save_text(self, text, absolute_filename, text_opt=''):
open(absolute_filename, "w").write(text)
return (text, )
class Wait_And_Read_File:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"absolute_filename": ("STRING", {
"default":
"image.txt",
"tooltip":
"The absolute path to the file to read."
})
},
"optional": {
"text": ("STRING", {
"multiline":
True,
"dynamicPrompts":
True,
"tooltip":
"Text to provide order when necessary (to wait on done file)."
}),
}
}
RETURN_TYPES = ("STRING", )
RETURN_NAMES = ("text from file", )
FUNCTION = "read_text"
OUTPUT_NODE = True
CATEGORY = "text"
DESCRIPTION = "Saves the input text to an absolute path."
def read_text(self, absolute_filename, text=''):
while not os.path.exists(absolute_filename):
time.sleep(0.1)
res = open(absolute_filename, "r").read()
os.unlink(absolute_filename)
return (res, )
+166
View File
@@ -0,0 +1,166 @@
from __future__ import annotations
from weakref import ref as WeakRef
from pathlib import Path
from tqdm import tqdm
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel
from kornia.filters import box_blur
def cosine_similarity(image_embeds: Tensor, text_embeds: Tensor):
if image_embeds.dim() == 2 and text_embeds.dim() == 2:
image_embeds = image_embeds.unsqueeze(1)
return F.cosine_similarity(image_embeds, text_embeds, dim=-1)
class CLIPSafetyChecker(PreTrainedModel):
# https://huggingface.co/CompVis/stable-diffusion-safety-checker
# Adapted from:
# https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/stable_diffusion/safety_checker.py
config_class = CLIPConfig
_no_split_modules = ["CLIPEncoderLayer"]
def __init__(self, config: CLIPConfig):
super().__init__(config)
projdim = config.projection_dim
self.vision_model = CLIPVisionModel(config.vision_config)
self.visual_projection = nn.Linear(config.vision_config.hidden_size, projdim, bias=False)
self.concept_embeds = nn.Parameter(torch.ones(17, projdim), requires_grad=False)
self.special_care_embeds = nn.Parameter(torch.ones(3, projdim), requires_grad=False)
self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False)
self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False)
def forward(self, clip_input, images: Tensor, sensitivity: float, alternate_image: Tensor):
with torch.no_grad():
image_batch = self.vision_model(clip_input)[1]
image_embeds = self.visual_projection(image_batch)
sensitivity = -0.1 + 0.14 * sensitivity
special_cos_dist = cosine_similarity(image_embeds, self.special_care_embeds)
special_scores_threshold = self.special_care_embeds_weights.unsqueeze(0)
special_scores = special_cos_dist - special_scores_threshold + sensitivity
if torch.any(special_scores > 0):
sensitivity = sensitivity + 0.01
cos_dist = cosine_similarity(image_embeds, self.concept_embeds)
concept_threshold = self.concept_embeds_weights.unsqueeze(0)
concept_scores = cos_dist - concept_threshold + sensitivity
is_nsfw = [torch.any(concept_scores[i] > 0) for i in range(concept_scores.shape[0])]
is_nsfw = [x.item() for x in is_nsfw]
return self.filter_images(images, alternate_image, is_nsfw)
def filter_images(self, images: Tensor, alternate_image: Tensor, is_nsfw: list[bool]):
if not any(is_nsfw):
return images
images = images.clone()
for idx, nsfw in enumerate(is_nsfw):
if nsfw:
# Resize alternate image to match original image dimensions
resized_alternate = F.interpolate(
alternate_image[idx:idx+1], # Add batch dimension
size=(images[idx].shape[1], images[idx].shape[2]), # Target height, width
mode='bilinear',
align_corners=False
)
images[idx] = resized_alternate.squeeze(0) # Remove batch dimension
return images
class CachedModels:
_instance: WeakRef | None = None
def __init__(self):
model_dir = Path(__file__).parent / "safetychecker"
model_file = model_dir / "model.safetensors"
if not model_file.exists():
self.download(
"https://huggingface.co/CompVis/stable-diffusion-safety-checker/resolve/refs%2Fpr%2F41/model.safetensors",
target=model_file,
)
self.feature_extractor = CLIPImageProcessor.from_pretrained(model_dir)
self.safety_checker = CLIPSafetyChecker.from_pretrained(model_dir)
@classmethod
def load(cls):
models = cls._instance and cls._instance()
if models is None:
models = cls()
cls._instance = WeakRef(models)
return models
def download(self, url: str, target: Path):
import requests
try:
target_temp = target.with_suffix(".download")
with requests.get(url, stream=True) as response:
text = "NSFWFilter model download"
total = int(response.headers.get("content-length", 0))
pbar = tqdm(None, total=total, unit="b", unit_scale=True, desc=text)
with open(target_temp, "wb") as f:
for chunk in response.iter_content(chunk_size=8192):
f.write(chunk)
pbar.update(len(chunk))
pbar.close()
target_temp.rename(target)
except Exception as e:
raise RuntimeError(
f"NSFWFilter: Failed to download safety-checker model from {url} to target location {target}: {e}"
) from e
def to_bchw(image: torch.Tensor):
if image.ndim == 3:
image = image.unsqueeze(0)
return image.movedim(-1, 1)
def to_bhwc(image: torch.Tensor):
return image.movedim(1, -1)
def mask_batch(mask: torch.Tensor):
if mask.ndim == 2:
mask = mask.unsqueeze(0)
return mask
class TRI3DNSFWFilter:
models: CachedModels
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"image": ("IMAGE",),
"alternate_image": ("IMAGE",),
"sensitivity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.10}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "check"
CATEGORY = "TRI3D NSFW"
def __init__(self):
self.models = CachedModels.load()
def check(self, image, alternate_image,sensitivity):
image = to_bchw(image)
alternate_image = to_bchw(alternate_image)
input = self.models.feature_extractor(image, do_rescale=False, return_tensors="pt")
filtered = self.models.safety_checker(
images=image, clip_input=input.pixel_values, sensitivity=sensitivity, alternate_image=alternate_image
)
return (to_bhwc(filtered),)
+171
View File
@@ -0,0 +1,171 @@
{
"_name_or_path": "clip-vit-large-patch14/",
"architectures": [
"SafetyChecker"
],
"initializer_factor": 1.0,
"logit_scale_init_value": 2.6592,
"model_type": "clip",
"projection_dim": 768,
"text_config": {
"_name_or_path": "",
"add_cross_attention": false,
"architectures": null,
"attention_dropout": 0.0,
"bad_words_ids": null,
"bos_token_id": 0,
"chunk_size_feed_forward": 0,
"cross_attention_hidden_size": null,
"decoder_start_token_id": null,
"diversity_penalty": 0.0,
"do_sample": false,
"dropout": 0.0,
"early_stopping": false,
"encoder_no_repeat_ngram_size": 0,
"eos_token_id": 2,
"exponential_decay_length_penalty": null,
"finetuning_task": null,
"forced_bos_token_id": null,
"forced_eos_token_id": null,
"hidden_act": "quick_gelu",
"hidden_size": 768,
"id2label": {
"0": "LABEL_0",
"1": "LABEL_1"
},
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 3072,
"is_decoder": false,
"is_encoder_decoder": false,
"label2id": {
"LABEL_0": 0,
"LABEL_1": 1
},
"layer_norm_eps": 1e-05,
"length_penalty": 1.0,
"max_length": 20,
"max_position_embeddings": 77,
"min_length": 0,
"model_type": "clip_text_model",
"no_repeat_ngram_size": 0,
"num_attention_heads": 12,
"num_beam_groups": 1,
"num_beams": 1,
"num_hidden_layers": 12,
"num_return_sequences": 1,
"output_attentions": false,
"output_hidden_states": false,
"output_scores": false,
"pad_token_id": 1,
"prefix": null,
"problem_type": null,
"pruned_heads": {},
"remove_invalid_values": false,
"repetition_penalty": 1.0,
"return_dict": true,
"return_dict_in_generate": false,
"sep_token_id": null,
"task_specific_params": null,
"temperature": 1.0,
"tie_encoder_decoder": false,
"tie_word_embeddings": true,
"tokenizer_class": null,
"top_k": 50,
"top_p": 1.0,
"torch_dtype": null,
"torchscript": false,
"transformers_version": "4.21.0.dev0",
"typical_p": 1.0,
"use_bfloat16": false,
"vocab_size": 49408
},
"text_config_dict": {
"hidden_size": 768,
"intermediate_size": 3072,
"num_attention_heads": 12,
"num_hidden_layers": 12
},
"torch_dtype": "float32",
"transformers_version": null,
"vision_config": {
"_name_or_path": "",
"add_cross_attention": false,
"architectures": null,
"attention_dropout": 0.0,
"bad_words_ids": null,
"bos_token_id": null,
"chunk_size_feed_forward": 0,
"cross_attention_hidden_size": null,
"decoder_start_token_id": null,
"diversity_penalty": 0.0,
"do_sample": false,
"dropout": 0.0,
"early_stopping": false,
"encoder_no_repeat_ngram_size": 0,
"eos_token_id": null,
"exponential_decay_length_penalty": null,
"finetuning_task": null,
"forced_bos_token_id": null,
"forced_eos_token_id": null,
"hidden_act": "quick_gelu",
"hidden_size": 1024,
"id2label": {
"0": "LABEL_0",
"1": "LABEL_1"
},
"image_size": 224,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 4096,
"is_decoder": false,
"is_encoder_decoder": false,
"label2id": {
"LABEL_0": 0,
"LABEL_1": 1
},
"layer_norm_eps": 1e-05,
"length_penalty": 1.0,
"max_length": 20,
"min_length": 0,
"model_type": "clip_vision_model",
"no_repeat_ngram_size": 0,
"num_attention_heads": 16,
"num_beam_groups": 1,
"num_beams": 1,
"num_hidden_layers": 24,
"num_return_sequences": 1,
"output_attentions": false,
"output_hidden_states": false,
"output_scores": false,
"pad_token_id": null,
"patch_size": 14,
"prefix": null,
"problem_type": null,
"pruned_heads": {},
"remove_invalid_values": false,
"repetition_penalty": 1.0,
"return_dict": true,
"return_dict_in_generate": false,
"sep_token_id": null,
"task_specific_params": null,
"temperature": 1.0,
"tie_encoder_decoder": false,
"tie_word_embeddings": true,
"tokenizer_class": null,
"top_k": 50,
"top_p": 1.0,
"torch_dtype": null,
"torchscript": false,
"transformers_version": "4.21.0.dev0",
"typical_p": 1.0,
"use_bfloat16": false
},
"vision_config_dict": {
"hidden_size": 1024,
"intermediate_size": 4096,
"num_attention_heads": 16,
"num_hidden_layers": 24,
"patch_size": 14
}
}
+20
View File
@@ -0,0 +1,20 @@
{
"crop_size": 224,
"do_center_crop": true,
"do_convert_rgb": true,
"do_normalize": true,
"do_resize": true,
"feature_extractor_type": "CLIPFeatureExtractor",
"image_mean": [
0.48145466,
0.4578275,
0.40821073
],
"image_std": [
0.26862954,
0.26130258,
0.27577711
],
"resample": 3,
"size": 224
}
+820
View File
@@ -0,0 +1,820 @@
import numpy as np
import torch
import json
import cv2
# {0, "Nose"},
# // {1, "Neck"},
# // {2, "RShoulder"},
# // {3, "RElbow"},
# // {4, "RWrist"},
# // {5, "LShoulder"},
# // {6, "LElbow"},
# // {7, "LWrist"},
# // {8, "MidHip"},
# // {9, "RHip"},
# // {10, "RKnee"},
# // {11, "RAnkle"},
# // {12, "LHip"},
# // {13, "LKnee"},
# // {14, "LAnkle"},
# // {15, "REye"},
# // {16, "LEye"},
# // {17, "REar"},
# // {18, "LEar"},
# // {19, "LBigToe"},
# // {20, "LSmallToe"},
# // {21, "LHeel"},
# // {22, "RBigToe"},
# // {23, "RSmallToe"},
# // {24, "RHeel"},
# // {25, "Background"}
class TRI3D_SmartBox:
def from_torch_image(self, image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(self, image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"keypoints_json": ("STRING", {"multiline": True}),
},
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", )
CATEGORY = "TRI3D"
def extract_torso_keypoints(self, keypoints):
# Indices for torso-related keypoints
torso_indices = [8, 9, 10, 11, 12, 13]
return [keypoints[i] for i in torso_indices]
def run(self, image, keypoints_json):
kp_data = json.loads(open(keypoints_json, 'r').read())
original_height, original_width = kp_data['height'], kp_data['width']
torso_keypoints = self.extract_torso_keypoints(kp_data['keypoints'])
# Convert Torch image to OpenCV format
cv_image = self.from_torch_image(image)
# Remove the batch dimension if present
if len(cv_image.shape) == 4:
cv_image = cv_image[0]
# Adjust keypoints to match the image dimensions
adjusted_keypoints = self.adjust_keypoints(torso_keypoints, cv_image.shape, original_height, original_width)
# Fill the area below the hip line
filled_image = self.fill_below_hip(cv_image, adjusted_keypoints)
# Convert back to Torch format
torch_image = self.to_torch_image(filled_image)
# Add the batch dimension back
torch_image = torch_image.unsqueeze(0)
return (torch_image,)
def adjust_keypoints(self, keypoints, image_shape, original_height, original_width):
image_height, image_width = image_shape[:2]
scale_x = image_width / original_width
scale_y = image_height / original_height
adjusted_keypoints = [
(int(x * scale_x), int(y * scale_y)) for x, y in keypoints
]
return adjusted_keypoints
def fill_below_hip(self, image, keypoints):
# Correct the indices for hip keypoints
# Assuming indices 8 and 11 are for left and right hips
# print(keypoints,"hip keypoints")
try:
valid_y_coords = [kp[1] for kp in keypoints if kp[1] >= 0]
hip_y = min(valid_y_coords) if valid_y_coords else 0
except:
hip_y = 0
if hip_y == 0:
return image
# Find the bounding box of the mask below the hip line
mask = image[:, :, 0] # Assuming single-channel mask
below_hip = mask[hip_y:, :]
contours, _ = cv2.findContours(below_hip, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cnt = 0
for contour in contours:
x, y, w, h = cv2.boundingRect(contour)
# print(cnt, x,y,w,h, cv2.contourArea(contour), "cnt,x,y,w,h,area")
cnt+=1
# cv2.rectangle(image, (x, y + hip_y), (x + w, y + h + hip_y), (255, 255, 255), -1)
contours = [contour for contour in contours if cv2.contourArea(contour) > 20]
if len(contours) == 0:
return image
# Combine all contours into one
all_contours = np.vstack(contours)
# Calculate a single bounding rectangle for all contours
x, y, w, h = cv2.boundingRect(all_contours)
# print(x,y,w,h, "x,y,w,h")
cv2.rectangle(image, (x, y + hip_y), (x + w, y + h + hip_y), (255, 255, 255), -1)
return image
class TRI3D_Skip_HeadMask:
def from_torch_image(self, image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(self, image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"head_mask": ("IMAGE", ),
},
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", )
CATEGORY = "TRI3D"
def run(self, image, head_mask):
# Convert Torch images to OpenCV format
cv_image = self.from_torch_image(image)
cv_head_mask = self.from_torch_image(head_mask)
# Remove the batch dimension if present
if len(cv_image.shape) == 4:
cv_image = cv_image[0]
if len(cv_head_mask.shape) == 4:
cv_head_mask = cv_head_mask[0]
# Find the lowest point in the head mask
mask = cv_head_mask[:, :, 0] # Assuming single-channel mask
contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
lowest_y = 0
for contour in contours:
for point in contour:
x, y = point[0]
if y > lowest_y:
lowest_y = y
# Black out everything above the lowest point
cv_image[:lowest_y, :] = 0
# Convert back to Torch format
torch_image = self.to_torch_image(cv_image)
# Add the batch dimension back
torch_image = torch_image.unsqueeze(0)
return (torch_image,)
class TRI3D_Skip_HeadMask_AddNeck:
def adjust_keypoints(self, keypoints, image_shape, original_height, original_width):
image_height, image_width = image_shape[:2]
scale_x = image_width / original_width
scale_y = image_height / original_height
adjusted_keypoints = [
(int(x * scale_x), int(y * scale_y)) for x, y in keypoints
]
return adjusted_keypoints
def from_torch_image(self, image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(self, image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def extract_neck_keypoint(self, keypoints):
# Indices for torso-related keypoints
neck_indices = [1]
return [keypoints[i] for i in neck_indices]
def extract_ear_keypoints(self, keypoints):
# Indices for ear keypoints (17=right ear, 18=left ear)
ear_indices = [17, 18]
return [keypoints[i] for i in ear_indices]
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"head_mask": ("IMAGE", ),
"keypoints_json": ("STRING", {"multiline": True}),
"ratio_aggression": ("FLOAT", {"default": 0.5, "min": 0, "max": 1, "step": 0.01}),
"neck_width_factor": ("FLOAT", {"default": 0.8, "min": 0.1, "max": 1.5, "step": 0.05}),
},
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", )
CATEGORY = "TRI3D"
def run(self, image, head_mask, keypoints_json, ratio_aggression, neck_width_factor):
# Convert Torch images to OpenCV format
cv_image = self.from_torch_image(image)
cv_head_mask = self.from_torch_image(head_mask)
# Remove the batch dimension if present
if len(cv_image.shape) == 4:
cv_image = cv_image[0]
if len(cv_head_mask.shape) == 4:
cv_head_mask = cv_head_mask[0]
kp_data = json.loads(open(keypoints_json, 'r').read())
original_height, original_width = kp_data['height'], kp_data['width']
neck_keypoints = self.extract_neck_keypoint(kp_data['keypoints'])
# Make a copy of the original image
result_image = cv_image.copy()
# Adjust keypoints to match the image dimensions
adjusted_neck_keypoints = self.adjust_keypoints(neck_keypoints, cv_image.shape, original_height, original_width)
# Find the lowest point and face dimensions in the head mask
mask = cv_head_mask[:, :, 0] # Assuming single-channel mask
contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
# Find the chin point (lowest point) and calculate face properties
lowest_y = 0
face_center_x = cv_image.shape[1] // 2 # Default to center of image
face_width = cv_image.shape[1] // 3 # Default face width
if contours:
# Find the lowest point (chin)
for contour in contours:
for point in contour:
x, y = point[0]
if y > lowest_y:
lowest_y = y
# Calculate face bounding box and center of gravity
x, y, w, h = cv2.boundingRect(contours[0])
face_width = w
# Calculate center of gravity of the face mask
M = cv2.moments(contours[0])
if M["m00"] != 0:
face_center_x = int(M["m10"] / M["m00"])
else:
face_center_x = x + w // 2
# Calculate weighted average point between neck and chin
neck_y = adjusted_neck_keypoints[0][1]
if neck_y <= 0:
neck_y = lowest_y
average_y = int((neck_y * ratio_aggression + lowest_y * (1 - ratio_aggression)))
print(neck_y, lowest_y, "neck_y, lowest_y")
print(average_y, "average_y")
# ZONE 1: Black out everything above the chin point
result_image[:lowest_y, :] = 0
# ZONE 2: Create a triangle for the neck area
if lowest_y < average_y: # Only process if there's a gap between chin and average_y
# Create a mask for Zone 2
zone2_mask = np.zeros_like(cv_image[:,:,0])
# Create a triangle with apex at weighted average point and base at chin level
# Apply the neck width factor to the face width
neck_width = int(face_width * neck_width_factor)
triangle_half_width = neck_width // 2
# Create polygon points for the triangle
triangle_points = np.array([
[face_center_x, average_y], # Apex at weighted average point
[face_center_x - triangle_half_width, lowest_y], # Left base point at chin level
[face_center_x + triangle_half_width, lowest_y] # Right base point at chin level
], dtype=np.int32)
# Fill the triangle in the mask
cv2.fillPoly(zone2_mask, [triangle_points], 255)
# Apply the mask only to the region between chin and weighted average
for y in range(lowest_y, average_y):
for x in range(cv_image.shape[1]):
if zone2_mask[y, x] > 0:
result_image[y, x] = 0
# ZONE 3: Area below weighted average point is left as is
# No action needed for this zone
# Convert back to Torch format
torch_image = self.to_torch_image(result_image)
# Add the batch dimension back
torch_image = torch_image.unsqueeze(0)
return (torch_image,)
class TRI3D_Image_extend:
def from_torch_image(self, image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(self, image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"face_mask": ("IMAGE", ),
"image": ("IMAGE", ),
"ratio": ("FLOAT", {"default": 1.5, "min": 1.2, "max": 2, "step": 0.01}),
},
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", "IMAGE", )
RETURN_NAMES = ("image", "mask_image", )
CATEGORY = "TRI3D"
def run(self, face_mask, image, ratio):
cv_face_mask = self.from_torch_image(face_mask)
cv_image = self.from_torch_image(image)
# Remove the batch dimension if present
if len(cv_image.shape) == 4:
cv_image = cv_image[0]
if len(cv_face_mask.shape) == 4:
cv_face_mask = cv_face_mask[0]
mask = cv_face_mask[:, :, 0] # Assuming single-channel mask
contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
lowest_y = 0
highest_y = cv_image.shape[0]
for contour in contours:
for point in contour:
x, y = point[0]
if y > lowest_y:
lowest_y = y
if y < highest_y:
highest_y = y
y_below_face = cv_image.shape[0] - lowest_y
y_face = lowest_y-highest_y
# Only extend if the space below face is less than 1.5 times face height
target_below_face = int(y_face * ratio)
# print("y_face", y_face)
# print("lowest_y", lowest_y)
# print("highest_y", highest_y)
# print("target_below_face", target_below_face)
# print("y_below_face", y_below_face)
original_height = cv_image.shape[0]
original_width = cv_image.shape[1]
if y_below_face < target_below_face:
y_extend = target_below_face - y_below_face
# Calculate how much to extend horizontally to maintain aspect ratio
new_height = original_height + y_extend
new_width = int(original_width * (new_height / original_height))
x_extend = new_width - original_width
x_extend_left = x_extend // 2
x_extend_right = x_extend - x_extend_left
# Extend the image in all necessary directions
cv_image = cv2.copyMakeBorder(
cv_image,
0, y_extend, # top, bottom
x_extend_left, x_extend_right, # left, right
cv2.BORDER_CONSTANT,
value=[0, 0, 0]
)
# Create extension mask
extension_mask = np.zeros_like(cv_image)
# Make extended portions white
extension_mask[original_height:, :] = 255 # bottom extension
extension_mask[:, :x_extend_left] = 255 # left extension
extension_mask[:, -x_extend_right:] = 255 # right extension
else:
extension_mask = np.zeros_like(cv_image)
# Convert both images back to torch format
torch_image = self.to_torch_image(cv_image)
torch_mask = self.to_torch_image(extension_mask)
# Add batch dimension to both
torch_image = torch_image.unsqueeze(0)
torch_mask = torch_mask.unsqueeze(0)
return (torch_image, torch_mask)
class TRI3D_Smart_Depth:
def from_torch_image(self, image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(self, image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"keypoints_json": ("STRING", {"multiline": True}),
},
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", )
CATEGORY = "TRI3D"
def extract_torso_keypoints(self, keypoints):
# Indices for torso-related keypoints
torso_indices = [8, 9, 10, 11, 12, 13]
return [keypoints[i] for i in torso_indices]
def run(self, image, keypoints_json):
kp_data = json.loads(open(keypoints_json, 'r').read())
original_height, original_width = kp_data['height'], kp_data['width']
torso_keypoints = self.extract_torso_keypoints(kp_data['keypoints'])
# Convert Torch image to OpenCV format
cv_image = self.from_torch_image(image)
# Remove the batch dimension if present
if len(cv_image.shape) == 4:
cv_image = cv_image[0]
# Adjust keypoints to match the image dimensions
adjusted_keypoints = self.adjust_keypoints(torso_keypoints, cv_image.shape, original_height, original_width)
# Fill the area below the hip line
filled_image = self.fill_below_hip(cv_image, adjusted_keypoints)
# Convert back to Torch format
torch_image = self.to_torch_image(filled_image)
# Add the batch dimension back
torch_image = torch_image.unsqueeze(0)
return (torch_image,)
def adjust_keypoints(self, keypoints, image_shape, original_height, original_width):
image_height, image_width = image_shape[:2]
scale_x = image_width / original_width
scale_y = image_height / original_height
adjusted_keypoints = [
(int(x * scale_x), int(y * scale_y)) for x, y in keypoints
]
return adjusted_keypoints
def fill_below_hip(self, image, keypoints):
# Correct the indices for hip keypoints
# Assuming indices 8 and 11 are for left and right hips
# print(keypoints,"hip keypoints")
try:
valid_y_coords = [kp[1] for kp in keypoints if kp[1] >= 0]
hip_y = min(valid_y_coords) if valid_y_coords else 0
except:
hip_y = 0
if hip_y == 0:
return image
# Find the bounding box of the mask below the hip line
mask = image[:, :, 0] # Assuming single-channel mask
below_hip = mask[hip_y:, :]
contours, _ = cv2.findContours(below_hip, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cnt = 0
for contour in contours:
x, y, w, h = cv2.boundingRect(contour)
# print(cnt, x,y,w,h, cv2.contourArea(contour), "cnt,x,y,w,h,area")
cnt+=1
# cv2.rectangle(image, (x, y + hip_y), (x + w, y + h + hip_y), (255, 255, 255), -1)
contours = [contour for contour in contours if cv2.contourArea(contour) > 0]
if len(contours) == 0:
return image
# Combine all contours into one
all_contours = np.vstack(contours)
# Calculate a single bounding rectangle for all contours
x, y, w, h = cv2.boundingRect(all_contours)
# print(x,y,w,h, "x,y,w,h")
cv2.rectangle(image, (x, y + hip_y), (x + w, y + h + hip_y), (0, 0, 0), -1)
return image
class TRI3D_NarrowfyImage:
def from_torch_image(self, image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(self, image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"mask": ("IMAGE", ),
"aspect_ratio": ("FLOAT", {"default": 0.33, "min": 0.25, "max": 1, "step": 0.01}),
"border_margin": ("INT", {"default": 15, "min": 10, "max": 100, "step": 1}),
},
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", "IMAGE", "INT", "INT",)
RETURN_NAMES = ("cropped_image", "cropped_mask", "cropped_width", "cropped_height",)
CATEGORY = "TRI3D"
def run(self, image, mask, aspect_ratio, border_margin):
# Convert to CV format and remove batch dimension
cv_image = self.from_torch_image(image)[0]
cv_mask = self.from_torch_image(mask)[0]
# Find bounding box of the mask
mask_channel = cv_mask[:, :, 0]
contours, _ = cv2.findContours(mask_channel, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
return image, mask, aspect_ratio
# Filter contours by area
significant_contours = [cnt for cnt in contours if cv2.contourArea(cnt) > 100]
if not significant_contours:
return image, mask, aspect_ratio
# Get combined bounding box for all significant contours
x_min = float('inf')
y_min = float('inf')
x_max = 0
y_max = 0
for contour in significant_contours:
x, y, w, h = cv2.boundingRect(contour)
x_min = min(x_min, x)
y_min = min(y_min, y)
x_max = max(x_max, x + w)
y_max = max(y_max, y + h)
# Calculate final width and height with margin
margin = border_margin
x = max(0, x_min - margin) # Ensure we don't go below 0
y = max(0, y_min - margin)
w = min(cv_image.shape[1] - x, (x_max - x_min) + 2 * margin) # Ensure we don't exceed image width
h = min(cv_image.shape[0] - y, (y_max - y_min) + 2 * margin) # Ensure we don't exceed image height
# Crop both image and mask to bounding box
cropped_image = cv_image[y:y+h, x:x+w]
cropped_mask = cv_mask[y:y+h, x:x+w]
# Calculate required height for aspect ratio 1/3
min_height = w * 1/aspect_ratio
if h < min_height:
height_extend = min_height - h
# Extend image with black pixels
extended_image = cv2.copyMakeBorder(
cropped_image,
0, int(height_extend), # top, bottom
0, 0, # left, right
cv2.BORDER_CONSTANT,
value=[0, 0, 0]
)
# Create mask with white pixels only in extended region
extended_mask = cv2.copyMakeBorder(
np.zeros_like(cropped_mask), # Start with black base
0, int(height_extend), # top, bottom
0, 0, # left, right
cv2.BORDER_CONSTANT,
value=[255, 255, 255] # White extension
)
cropped_image = extended_image
cropped_mask = extended_mask
# Convert back to torch format and add batch dimension
torch_image = self.to_torch_image(cropped_image).unsqueeze(0)
torch_mask = self.to_torch_image(cropped_mask).unsqueeze(0)
return (torch_image, torch_mask,w,h)
class TRI3D_CropAndExtend:
def from_torch_image(self, image):
image = image.cpu().numpy() * 255.0
image = np.clip(image, 0, 255).astype(np.uint8)
return image
def to_torch_image(self, image):
image = image.astype(dtype=np.float32)
image /= 255.0
image = torch.from_numpy(image)
return image
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"garment_image": ("IMAGE",),
"garment_mask": ("IMAGE",),
"human_image": ("IMAGE",),
"human_mask": ("IMAGE",),
"margin": ("INT", {"default": 10, "min": 0, "max": 50}),
},
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", "IMAGE", "IMAGE", "IMAGE", "INT", "INT",)
RETURN_NAMES = ("cropped_garment", "cropped_garment_mask", "cropped_human", "cropped_human_mask", "cropped_width", "cropped_height",)
def run(self, garment_image, garment_mask, human_image, human_mask, margin):
# Convert to CV format and remove batch dimension
cv_garment = self.from_torch_image(garment_image)[0]
cv_garment_mask = self.from_torch_image(garment_mask)[0]
cv_human = self.from_torch_image(human_image)[0]
cv_human_mask = self.from_torch_image(human_mask)[0]
# Process garment
mask_channel = cv_garment_mask[:, :, 0]
contours, _ = cv2.findContours(mask_channel, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
return garment_image, garment_mask, human_image, human_mask, cv_garment.shape[1], cv_garment.shape[0]
# Get bounding box with margin
x, y, w, h = cv2.boundingRect(contours[0])
x = max(0, x - margin)
y = max(0, y - margin)
w = min(cv_garment.shape[1] - x, w + 2 * margin)
h = min(cv_garment.shape[0] - y, h + 2 * margin)
# Store the cropped dimensions before extension
cropped_width = w
cropped_height = h
# Crop garment and its mask
cropped_garment = cv_garment[y:y+h, x:x+w]
cropped_garment_mask = cv_garment_mask[y:y+h, x:x+w]
# Calculate required height for aspect ratio 1/3
min_height = w * 3
if h < min_height:
height_extend = min_height - h
# Extend garment image and mask
extended_garment = cv2.copyMakeBorder(
cropped_garment,
0, int(height_extend),
0, 0,
cv2.BORDER_CONSTANT,
value=[0, 0, 0]
)
extended_garment_mask = cv2.copyMakeBorder(
cropped_garment_mask,
0, int(height_extend),
0, 0,
cv2.BORDER_CONSTANT,
value=[255, 255, 255]
)
cropped_garment = extended_garment
cropped_garment_mask = extended_garment_mask
# Process human image similarly
mask_channel = cv_human_mask[:, :, 0]
contours, _ = cv2.findContours(mask_channel, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if contours:
x, y, w, h = cv2.boundingRect(contours[0])
x = max(0, x - margin)
y = max(0, y - margin)
w = min(cv_human.shape[1] - x, w + 2 * margin)
h = min(cv_human.shape[0] - y, h + 2 * margin)
cropped_human = cv_human[y:y+h, x:x+w]
cropped_human_mask = cv_human_mask[y:y+h, x:x+w]
min_height = w * 3
if h < min_height:
height_extend = min_height - h
extended_human = cv2.copyMakeBorder(
cropped_human,
0, int(height_extend),
0, 0,
cv2.BORDER_CONSTANT,
value=[0, 0, 0]
)
extended_human_mask = cv2.copyMakeBorder(
cropped_human_mask,
0, int(height_extend),
0, 0,
cv2.BORDER_CONSTANT,
value=[255, 255, 255]
)
cropped_human = extended_human
cropped_human_mask = extended_human_mask
# Convert back to torch format and add batch dimension
torch_garment = self.to_torch_image(cropped_garment).unsqueeze(0)
torch_garment_mask = self.to_torch_image(cropped_garment_mask).unsqueeze(0)
torch_human = self.to_torch_image(cropped_human).unsqueeze(0)
torch_human_mask = self.to_torch_image(cropped_human_mask).unsqueeze(0)
return (torch_garment, torch_garment_mask, torch_human, torch_human_mask, cropped_width, cropped_height)
+409
View File
@@ -0,0 +1,409 @@
import torch, cv2, json
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
class TRI3D_clean_mask():
"""For the given mask and threshold area, remove all patches in the mask with area smaller than threshold"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"masks": ("MASK", ),
"threshold":("FLOAT",{"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01})
}
}
FUNCTION = "run"
RETURN_TYPES = ("MASK", "BOOL")
RETURN_NAMES = ("mask", "cleaned")
CATEGORY = "TRI3D"
def run(self, masks, threshold):
batch_results = []
for mask in masks:
mask = from_torch_image(mask)
mask = np.where(mask < 127, 0, 255).astype(np.uint8)
h,w = mask.shape[:2]
total_area = h*w
# num_labels, labels = cv2.connectedComponents(mask)
region_mask = np.zeros_like(mask)
# for label in range(1, num_labels):
# area_percent = (np.sum(labels == label)/ total_area) * 100
# if area_percent < threshold:
# continue
# region_mask[labels == label] = 255
less_than_threshold = True
area_percent = (np.sum(mask == 255)/ total_area) * 100
if area_percent > threshold:
region_mask[mask == 255] = 255
less_than_threshold = False
region_mask = to_torch_image(region_mask)
batch_results.append(region_mask.squeeze(0))
batch_results = torch.stack(batch_results)
return (batch_results, less_than_threshold)
class TRI3D_extract_pose_part():
"""
For the given pose, extract region around body parts, region can be defined by % of image size
"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"pose_json": ("STRING",{"default" : "dwpose/keypoints/input.json"}),
"width_pad": ("FLOAT",{"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}),
"height_pad": ("FLOAT",{"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}),
"shoulders":("BOOLEAN", {
"default": False
})
}
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", "STRING")
RETURN_NAMES = ("image", "coords")
CATEGORY = "TRI3D"
def get_frame_coords(self,point1, point2):
x1, y1 = point1
x2, y2 = point2
xmin, xmax, ymin, ymax = min(x1, x2), max(x1, x2), min(y1, y2), max(y1, y2)
for i in [xmin, xmax, ymin, ymax]:
if i < 0:
return None
return [xmin, xmax, ymin, ymax]
def run(self, image, pose_json, width_pad, height_pad, shoulders):
"""
image : input image
width_pad: % of image width you want to apply on both size of pose body part
height_pad: % of image width you want to apply on both size of pose body part
rest of them are body parts
"""
image = from_torch_image(image[0])
batch_result = []
input_pose = json.load(open(pose_json))
keypoints = input_pose['keypoints']
og_h, og_w = image.shape[:2]
ph, pw = [input_pose['height'], input_pose['width']]
for i,point in enumerate(keypoints):
x,y = point
y = int((y/ph)*og_h)
x = int((x/pw)*og_w)
keypoints[i] = [x, y]
width_offset = int(og_w * (width_pad) / 100)
height_offset = int(og_h * (height_pad) / 100)
xmin, xmax, ymin, ymax = [0, og_w, 0, og_h]
part_to_coords = {
"shoulders":self.get_frame_coords(keypoints[2], keypoints[5])
}
if shoulders:
print(part_to_coords["shoulders"])
if part_to_coords["shoulders"] != None:
new_xmin, new_xmax, new_ymin, new_ymax = part_to_coords["shoulders"]
xmin, xmax, ymin, ymax = new_xmin, new_xmax, new_ymin, new_ymax
xmin = max(0, xmin - width_offset)
xmax = min(og_w, xmax + width_offset)
ymin = max(0, ymin - height_offset)
ymax = min(og_h, ymax + height_offset)
image = image[ymin:ymax, xmin:xmax, :].astype(np.uint8)
image = to_torch_image(image)
batch_result.append(image)
batch_result = torch.stack(batch_result)
print("final_coords", xmin, xmax, ymin, ymax)
coords = ",".join([str(xmin), str(xmax), str(ymin), str(ymax)])
return batch_result, coords
class TRI3D_position_pose_part():
"""
put back extracted parts on OG image
"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"og_image": ("IMAGE", ),
"extracted_image": ("IMAGE", ),
"coords": ("STRING",{"default" : "xmin, xmax, ymin, ymax"}),
}
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("image", )
CATEGORY = "TRI3D"
def run(self, og_image, extracted_image, coords):
batch_result = []
og_image = from_torch_image(og_image[0])
extracted_image = from_torch_image(extracted_image[0])
xmin, xmax, ymin, ymax = [int(i) for i in coords.split(",")]
og_image[ymin:ymax, xmin:xmax, :] = extracted_image
og_image = to_torch_image(og_image).unsqueeze(0)
batch_result.append(og_image)
batch_result = torch.stack(batch_result)
return batch_result
class TRI3D_fill_mask():
"""
fill mask with the neighbouring pixels
"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE", ),
"mask": ("MASK", ),
"negative_mask": ("MASK", ),
"offset":("FLOAT",{"default": 1, "min": 0.0, "max": 100.0, "step": 0.01})
}
}
FUNCTION = "run"
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
CATEGORY = "TRI3D"
def run(self, image, mask, negative_mask, offset):
image = from_torch_image(image[0])
mask = mask[0].cpu().numpy()
mask = np.expand_dims(mask, -1)
mh, mw, _ = mask.shape
inverse_mask = np.ones_like(mask) - mask
negative_mask = negative_mask[0].cpu().numpy()
indices = np.where(mask > 0)
offset = offset / 100
source = image.copy()
for y,x in zip(indices[0],indices[1]):
x_off = min(mw-1, int(x + offset * mw))
if negative_mask[y][x_off] == 0: #check if pixles on right are outside body
source[y][x] = image[y][x_off]
else:
x_off = max(0, int(x - offset * mw)) #check if pixles on left are outside body
if negative_mask[y][x_off] == 0:
source[y][x] = image[y][x_off]
else:
y_off = max(0, int(y - offset * mh))
if negative_mask[y_off][x] == 0: #check if pixles on top are outside body
source[y][x] = image[y_off][x]
else:
y_off = min(mh-1, int(y + offset * mh))
if negative_mask[y_off][x] == 0: #check if pixles on bottom are outside body
source[y][x] = image[y_off][x]
image = mask * source + inverse_mask * image
image = to_torch_image(image).unsqueeze(0)
return (image,)
class TRI3D_is_only_trouser:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"pose_json_file": ("STRING", {
"default": "dwpose/keypoints"
})
}
}
RETURN_TYPES = ("BOOLEAN", )
FUNCTION = "main"
CATEGORY = "TRI3D"
def main(self, pose_json_file):
pose = json.load(open(pose_json_file))
height = pose['height']
width = pose['width']
keypoints = pose['keypoints']
points = [0,14,15,16,17,2,1,5]
point_to_part = {0:'nose',14:"left eye",15:"right eye",16:"left ear",17:"right ear",2:"left shoulder",1:"neck",5:"right shoulder"}
all_negative = True #if all face and shoulder points are negative means it is a bottom shot
for point in points:
x,y = keypoints[point]
if x > 0 and y > 0:
all_negative = False
print(f"{point_to_part[point]} exist")
return (all_negative,)
class TRI3D_extract_facer_mask:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"background": ("BOOLEAN", {
"default": False
}),
'hair':("BOOLEAN", {
"default": False
}),
'lower_lip':("BOOLEAN", {
"default": False
}),
'inner_mouth':("BOOLEAN", {
"default": False
}),
'upper_lip':("BOOLEAN", {
"default": False
}),
'nose':("BOOLEAN", {
"default": False
}),
'left_eyebrow':("BOOLEAN", {
"default": False
}),
'right_eyebrow':("BOOLEAN", {
"default": False
}),
'left_eye':("BOOLEAN", {
"default": False
}),
'right_eye':("BOOLEAN", {
"default": False
}),
'face':("BOOLEAN", {
"default": False
})
}
}
RETURN_TYPES = ("MASK", )
FUNCTION = "main"
CATEGORY = "TRI3D"
def main(self, image, background, hair, lower_lip, inner_mouth, upper_lip, nose, left_eyebrow, right_eyebrow, left_eye, right_eye, face):
image = from_torch_image(image[0])
h,w,_ = image.shape
mask = np.zeros_like(image)
label_to_rgb = {'background':[0,0,0], 'face':[0,138,255], 'right_eye':[180, 255, 0], 'left_eye':[42, 255, 0], 'right_eyebrow':[0, 255, 96],
'left_eyebrow':[0,255,234], 'nose':[255, 192, 0], 'upper_lip':[255, 54, 0], 'inner_mouth':[255, 0, 84], 'lower_lip':[255, 0, 222],
'hair':[150,0,255]}
if background:
temp = np.all(image == label_to_rgb['background'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if face:
temp = np.all(image == label_to_rgb['face'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if right_eye:
temp = np.all(image == label_to_rgb['right_eye'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if left_eye:
temp = np.all(image == label_to_rgb['left_eye'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if right_eyebrow:
temp = np.all(image == label_to_rgb['right_eyebrow'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if left_eyebrow:
temp = np.all(image == label_to_rgb['left_eyebrow'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if nose:
temp = np.all(image == label_to_rgb['nose'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if upper_lip:
temp = np.all(image == label_to_rgb['upper_lip'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if inner_mouth:
temp = np.all(image == label_to_rgb['inner_mouth'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if lower_lip:
temp = np.all(image == label_to_rgb['lower_lip'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
if hair:
temp = np.all(image == label_to_rgb['hair'], axis=-1)
idcs = np.where(temp==True)
mask[idcs] = 255
mask = to_torch_image(mask[:,:,0]).unsqueeze(0)
return (mask,)