Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
230131d73d | ||
|
|
85e3756721 | ||
|
|
9ccce9efed | ||
|
|
2c1173254e | ||
|
|
abaf432b6f | ||
|
|
a2c78015f4 | ||
|
|
1ef087fe27 | ||
|
|
4d6a17949c | ||
|
|
6d77b7e212 | ||
|
|
9e601095c4 | ||
|
|
340dcfdebd | ||
|
|
6dbc032baf | ||
|
|
11cfb6210d | ||
|
|
893d53778e | ||
|
|
981fd56671 | ||
|
|
983a680720 | ||
|
|
641b9cef7a | ||
|
|
5c89d3325d | ||
|
|
59aa80b723 | ||
|
|
f061a43fcb | ||
|
|
b89b0251b2 | ||
|
|
f9987b82c3 | ||
|
|
1eaa4d0a0c | ||
|
|
081b2f7acc | ||
|
|
f12aa9888d | ||
|
|
59fff21026 | ||
|
|
9af4bc2923 | ||
|
|
43815bf7c6 | ||
|
|
9a02c21a6a | ||
|
|
477efcf83d | ||
|
|
2fa2f55d2b | ||
|
|
ce4ff65e15 | ||
|
|
2322e1aa37 | ||
|
|
741d265ca5 | ||
|
|
abf527931b | ||
|
|
2a0be7fbb9 | ||
|
|
648e3cda90 | ||
|
|
4ccb25912e | ||
|
|
142749e496 | ||
|
|
749c7cfb6e | ||
|
|
dc91c7df3c | ||
|
|
5681fd0949 | ||
|
|
97e8a27e3b | ||
|
|
53d9ce4aac | ||
|
|
8528408548 | ||
|
|
b8459c7ae2 | ||
|
|
e0107c09e8 | ||
|
|
40f645f299 | ||
|
|
2a0e2b6191 | ||
|
|
7d0315bf11 | ||
|
|
deebf0f9ca | ||
|
|
7d419d7c92 | ||
|
|
9d15ec458e | ||
|
|
ebfd510ae3 | ||
|
|
c8f3ca0f29 | ||
|
|
c335e3dd8c | ||
|
|
ebd4cda04e | ||
|
|
4d587daa91 | ||
|
|
ff441da051 | ||
|
|
d8fd71bf4d | ||
|
|
3f887118d8 | ||
|
|
4b1aa9f3fa | ||
|
|
a63dbb238a | ||
|
|
4f8408af9c | ||
|
|
a7810c118b | ||
|
|
3a592a5376 | ||
|
|
e3c94acdfa | ||
|
|
5fd4064b1e | ||
|
|
7595071a3c | ||
|
|
7ccb3bed5d | ||
|
|
b723d3787f | ||
|
|
e14e6b378b | ||
|
|
dad3bc5bee | ||
|
|
a022ce7e01 | ||
|
|
c2d6abbbbe | ||
|
|
a3c16609a1 | ||
|
|
c1b47aa0d1 | ||
|
|
41ddedb0ba | ||
|
|
16247912db | ||
|
|
723eb31c79 | ||
|
|
2a0c9498a8 | ||
|
|
36684f5c29 | ||
|
|
204f68fddd | ||
|
|
5ab990ade6 | ||
|
|
afe299d36d | ||
|
|
e4410c53ed | ||
|
|
71415ae250 | ||
|
|
d0af6f652c | ||
|
|
1903728517 | ||
|
|
e922c37579 | ||
|
|
6e77293a3c | ||
|
|
6475c41e4a | ||
|
|
a2ab6c4945 | ||
|
|
a288558221 | ||
|
|
cae4239d18 | ||
|
|
4d91409365 | ||
|
|
e0e86cec0a | ||
|
|
bda29f6ba2 | ||
|
|
cb4103a167 | ||
|
|
fead250f55 | ||
|
|
1550bdec7d | ||
|
|
17b5a12fda | ||
|
|
7c0b9b9828 | ||
|
|
ac37c4c6ad | ||
|
|
3f2ee1d80a | ||
|
|
fd89979d62 | ||
|
|
ab9ab2b662 | ||
|
|
3702c0f0da | ||
|
|
4aa848028a | ||
|
|
8ba3531ab1 | ||
|
|
e6a8626537 | ||
|
|
ce528ac9c4 | ||
|
|
62350a2ff8 | ||
|
|
344bb9ebf7 |
@@ -1 +1,3 @@
|
||||
CLIPDROP_API_KEY=
|
||||
CLIPDROP_API_KEY=
|
||||
COMFY_PYTHON_PATH=/home/ubuntu/.conda/envs/comfy/bin/python
|
||||
PHOTOROOM_API_KEY=3603b83dfa1846bc3c7270ead7876
|
||||
@@ -6,5 +6,8 @@ venv
|
||||
checkpoints/
|
||||
checkpoint/
|
||||
.env
|
||||
.pth
|
||||
cloth-segmentation/model/cloth_segm.pth
|
||||
|
||||
dwpose/keypoints/
|
||||
dwpose/keypoints/
|
||||
huggingface/
|
||||
|
||||
@@ -88,9 +88,29 @@ def get_palette(num_cls):
|
||||
return palette
|
||||
|
||||
|
||||
def download_model_restore(model_restore_path):
|
||||
import os
|
||||
import gdown
|
||||
from pathlib import Path
|
||||
|
||||
# Ensure the directory for the model path exists
|
||||
os.makedirs(os.path.dirname(model_restore_path), exist_ok=True)
|
||||
|
||||
# Check if the model file already exists
|
||||
if not Path(model_restore_path).is_file():
|
||||
print("Model file does not exist, downloading...")
|
||||
|
||||
# Google Drive ID for the file
|
||||
# file_id = '1ruJg4lqR_jgQPj-9K0PP-L2vJERYOxLP'
|
||||
file_id="1AVVLm1LxOs3W1Fp_GLefIz6fdWfEYg88"
|
||||
gdown.download(id=file_id, output=model_restore_path, quiet=False)
|
||||
print("Download complete.")
|
||||
else:
|
||||
print("Model file already exists.")
|
||||
|
||||
|
||||
def main():
|
||||
args = get_arguments()
|
||||
|
||||
gpus = [int(i) for i in args.gpu.split(',')]
|
||||
assert len(gpus) == 1
|
||||
if not args.gpu == 'None':
|
||||
@@ -102,7 +122,7 @@ def main():
|
||||
print("Evaluating total class number {} with {}".format(num_classes, label))
|
||||
|
||||
model = networks.init_model('resnet101', num_classes=num_classes, pretrained=None)
|
||||
|
||||
download_model_restore(args.model_restore)
|
||||
state_dict = torch.load(args.model_restore)['state_dict']
|
||||
from collections import OrderedDict
|
||||
new_state_dict = OrderedDict()
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
import PIL
|
||||
import cv2
|
||||
import torch
|
||||
import os
|
||||
from process import load_seg_model, get_palette, generate_mask
|
||||
|
||||
device = 'cuda'
|
||||
|
||||
|
||||
def initialize_and_load_models():
|
||||
checkpoint_path = 'model/cloth_segm.pth'
|
||||
net = load_seg_model(checkpoint_path, device=device)
|
||||
return net
|
||||
|
||||
|
||||
net = initialize_and_load_models()
|
||||
|
||||
|
||||
def run(img):
|
||||
palette = get_palette(4)
|
||||
cloth_seg = generate_mask(img, net=net, device=device)
|
||||
return cloth_seg
|
||||
|
||||
|
||||
INPUT_PATH = "./input/"
|
||||
OUTPUT_PATH = "./output/"
|
||||
|
||||
import os
|
||||
for cur_image in os.listdir(INPUT_PATH):
|
||||
img = PIL.Image.open(INPUT_PATH + cur_image)
|
||||
cloth_seg = run(img)
|
||||
|
||||
cv2.imwrite(OUTPUT_PATH + cur_image,
|
||||
cv2.cvtColor(src=cloth_seg, code=cv2.COLOR_RGB2BGR))
|
||||
# cloth_seg.save(OUTPUT_PATH + cur_image, format="PNG")
|
||||
@@ -0,0 +1 @@
|
||||
/*upload model */
|
||||
@@ -0,0 +1,560 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class REBNCONV(nn.Module):
|
||||
def __init__(self, in_ch=3, out_ch=3, dirate=1):
|
||||
super(REBNCONV, self).__init__()
|
||||
|
||||
self.conv_s1 = nn.Conv2d(
|
||||
in_ch, out_ch, 3, padding=1 * dirate, dilation=1 * dirate
|
||||
)
|
||||
self.bn_s1 = nn.BatchNorm2d(out_ch)
|
||||
self.relu_s1 = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
xout = self.relu_s1(self.bn_s1(self.conv_s1(hx)))
|
||||
|
||||
return xout
|
||||
|
||||
|
||||
## upsample tensor 'src' to have the same spatial size with tensor 'tar'
|
||||
def _upsample_like(src, tar):
|
||||
|
||||
src = F.upsample(src, size=tar.shape[2:], mode="bilinear")
|
||||
|
||||
return src
|
||||
|
||||
|
||||
### RSU-7 ###
|
||||
class RSU7(nn.Module): # UNet07DRES(nn.Module):
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU7, self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool4 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool5 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv6 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
|
||||
self.rebnconv7 = REBNCONV(mid_ch, mid_ch, dirate=2)
|
||||
|
||||
self.rebnconv6d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv5d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv4d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
hx = self.pool4(hx4)
|
||||
|
||||
hx5 = self.rebnconv5(hx)
|
||||
hx = self.pool5(hx5)
|
||||
|
||||
hx6 = self.rebnconv6(hx)
|
||||
|
||||
hx7 = self.rebnconv7(hx6)
|
||||
|
||||
hx6d = self.rebnconv6d(torch.cat((hx7, hx6), 1))
|
||||
hx6dup = _upsample_like(hx6d, hx5)
|
||||
|
||||
hx5d = self.rebnconv5d(torch.cat((hx6dup, hx5), 1))
|
||||
hx5dup = _upsample_like(hx5d, hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5dup, hx4), 1))
|
||||
hx4dup = _upsample_like(hx4d, hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1))
|
||||
hx3dup = _upsample_like(hx3d, hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1))
|
||||
hx2dup = _upsample_like(hx2d, hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1))
|
||||
|
||||
"""
|
||||
del hx1, hx2, hx3, hx4, hx5, hx6, hx7
|
||||
del hx6d, hx5d, hx3d, hx2d
|
||||
del hx2dup, hx3dup, hx4dup, hx5dup, hx6dup
|
||||
"""
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
### RSU-6 ###
|
||||
class RSU6(nn.Module): # UNet06DRES(nn.Module):
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU6, self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool4 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
|
||||
self.rebnconv6 = REBNCONV(mid_ch, mid_ch, dirate=2)
|
||||
|
||||
self.rebnconv5d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv4d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
hx = self.pool4(hx4)
|
||||
|
||||
hx5 = self.rebnconv5(hx)
|
||||
|
||||
hx6 = self.rebnconv6(hx5)
|
||||
|
||||
hx5d = self.rebnconv5d(torch.cat((hx6, hx5), 1))
|
||||
hx5dup = _upsample_like(hx5d, hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5dup, hx4), 1))
|
||||
hx4dup = _upsample_like(hx4d, hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1))
|
||||
hx3dup = _upsample_like(hx3d, hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1))
|
||||
hx2dup = _upsample_like(hx2d, hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1))
|
||||
|
||||
"""
|
||||
del hx1, hx2, hx3, hx4, hx5, hx6
|
||||
del hx5d, hx4d, hx3d, hx2d
|
||||
del hx2dup, hx3dup, hx4dup, hx5dup
|
||||
"""
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
### RSU-5 ###
|
||||
class RSU5(nn.Module): # UNet05DRES(nn.Module):
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU5, self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=2)
|
||||
|
||||
self.rebnconv4d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
|
||||
hx5 = self.rebnconv5(hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5, hx4), 1))
|
||||
hx4dup = _upsample_like(hx4d, hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1))
|
||||
hx3dup = _upsample_like(hx3d, hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1))
|
||||
hx2dup = _upsample_like(hx2d, hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1))
|
||||
|
||||
"""
|
||||
del hx1, hx2, hx3, hx4, hx5
|
||||
del hx4d, hx3d, hx2d
|
||||
del hx2dup, hx3dup, hx4dup
|
||||
"""
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
### RSU-4 ###
|
||||
class RSU4(nn.Module): # UNet04DRES(nn.Module):
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU4, self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=2)
|
||||
|
||||
self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
|
||||
hx4 = self.rebnconv4(hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4, hx3), 1))
|
||||
hx3dup = _upsample_like(hx3d, hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1))
|
||||
hx2dup = _upsample_like(hx2d, hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1))
|
||||
|
||||
"""
|
||||
del hx1, hx2, hx3, hx4
|
||||
del hx3d, hx2d
|
||||
del hx2dup, hx3dup
|
||||
"""
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
### RSU-4F ###
|
||||
class RSU4F(nn.Module): # UNet04FRES(nn.Module):
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU4F, self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1)
|
||||
self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=2)
|
||||
self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=4)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=8)
|
||||
|
||||
self.rebnconv3d = REBNCONV(mid_ch * 2, mid_ch, dirate=4)
|
||||
self.rebnconv2d = REBNCONV(mid_ch * 2, mid_ch, dirate=2)
|
||||
self.rebnconv1d = REBNCONV(mid_ch * 2, out_ch, dirate=1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx2 = self.rebnconv2(hx1)
|
||||
hx3 = self.rebnconv3(hx2)
|
||||
|
||||
hx4 = self.rebnconv4(hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4, hx3), 1))
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3d, hx2), 1))
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2d, hx1), 1))
|
||||
|
||||
"""
|
||||
del hx1, hx2, hx3, hx4
|
||||
del hx3d, hx2d
|
||||
"""
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
##### U^2-Net ####
|
||||
class U2NET(nn.Module):
|
||||
def __init__(self, in_ch=3, out_ch=1):
|
||||
super(U2NET, self).__init__()
|
||||
|
||||
self.stage1 = RSU7(in_ch, 32, 64)
|
||||
self.pool12 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage2 = RSU6(64, 32, 128)
|
||||
self.pool23 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage3 = RSU5(128, 64, 256)
|
||||
self.pool34 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage4 = RSU4(256, 128, 512)
|
||||
self.pool45 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage5 = RSU4F(512, 256, 512)
|
||||
self.pool56 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage6 = RSU4F(512, 256, 512)
|
||||
|
||||
# decoder
|
||||
self.stage5d = RSU4F(1024, 256, 512)
|
||||
self.stage4d = RSU4(1024, 128, 256)
|
||||
self.stage3d = RSU5(512, 64, 128)
|
||||
self.stage2d = RSU6(256, 32, 64)
|
||||
self.stage1d = RSU7(128, 16, 64)
|
||||
|
||||
self.side1 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
self.side2 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
self.side3 = nn.Conv2d(128, out_ch, 3, padding=1)
|
||||
self.side4 = nn.Conv2d(256, out_ch, 3, padding=1)
|
||||
self.side5 = nn.Conv2d(512, out_ch, 3, padding=1)
|
||||
self.side6 = nn.Conv2d(512, out_ch, 3, padding=1)
|
||||
|
||||
self.outconv = nn.Conv2d(6 * out_ch, out_ch, 1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
|
||||
# stage 1
|
||||
hx1 = self.stage1(hx)
|
||||
hx = self.pool12(hx1)
|
||||
|
||||
# stage 2
|
||||
hx2 = self.stage2(hx)
|
||||
hx = self.pool23(hx2)
|
||||
|
||||
# stage 3
|
||||
hx3 = self.stage3(hx)
|
||||
hx = self.pool34(hx3)
|
||||
|
||||
# stage 4
|
||||
hx4 = self.stage4(hx)
|
||||
hx = self.pool45(hx4)
|
||||
|
||||
# stage 5
|
||||
hx5 = self.stage5(hx)
|
||||
hx = self.pool56(hx5)
|
||||
|
||||
# stage 6
|
||||
hx6 = self.stage6(hx)
|
||||
hx6up = _upsample_like(hx6, hx5)
|
||||
|
||||
# -------------------- decoder --------------------
|
||||
hx5d = self.stage5d(torch.cat((hx6up, hx5), 1))
|
||||
hx5dup = _upsample_like(hx5d, hx4)
|
||||
|
||||
hx4d = self.stage4d(torch.cat((hx5dup, hx4), 1))
|
||||
hx4dup = _upsample_like(hx4d, hx3)
|
||||
|
||||
hx3d = self.stage3d(torch.cat((hx4dup, hx3), 1))
|
||||
hx3dup = _upsample_like(hx3d, hx2)
|
||||
|
||||
hx2d = self.stage2d(torch.cat((hx3dup, hx2), 1))
|
||||
hx2dup = _upsample_like(hx2d, hx1)
|
||||
|
||||
hx1d = self.stage1d(torch.cat((hx2dup, hx1), 1))
|
||||
|
||||
# side output
|
||||
d1 = self.side1(hx1d)
|
||||
|
||||
d2 = self.side2(hx2d)
|
||||
d2 = _upsample_like(d2, d1)
|
||||
|
||||
d3 = self.side3(hx3d)
|
||||
d3 = _upsample_like(d3, d1)
|
||||
|
||||
d4 = self.side4(hx4d)
|
||||
d4 = _upsample_like(d4, d1)
|
||||
|
||||
d5 = self.side5(hx5d)
|
||||
d5 = _upsample_like(d5, d1)
|
||||
|
||||
d6 = self.side6(hx6)
|
||||
d6 = _upsample_like(d6, d1)
|
||||
|
||||
d0 = self.outconv(torch.cat((d1, d2, d3, d4, d5, d6), 1))
|
||||
|
||||
"""
|
||||
del hx1, hx2, hx3, hx4, hx5, hx6
|
||||
del hx5d, hx4d, hx3d, hx2d, hx1d
|
||||
del hx6up, hx5dup, hx4dup, hx3dup, hx2dup
|
||||
"""
|
||||
|
||||
return d0, d1, d2, d3, d4, d5, d6
|
||||
|
||||
|
||||
### U^2-Net small ###
|
||||
class U2NETP(nn.Module):
|
||||
def __init__(self, in_ch=3, out_ch=1):
|
||||
super(U2NETP, self).__init__()
|
||||
|
||||
self.stage1 = RSU7(in_ch, 16, 64)
|
||||
self.pool12 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage2 = RSU6(64, 16, 64)
|
||||
self.pool23 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage3 = RSU5(64, 16, 64)
|
||||
self.pool34 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage4 = RSU4(64, 16, 64)
|
||||
self.pool45 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage5 = RSU4F(64, 16, 64)
|
||||
self.pool56 = nn.MaxPool2d(2, stride=2, ceil_mode=True)
|
||||
|
||||
self.stage6 = RSU4F(64, 16, 64)
|
||||
|
||||
# decoder
|
||||
self.stage5d = RSU4F(128, 16, 64)
|
||||
self.stage4d = RSU4(128, 16, 64)
|
||||
self.stage3d = RSU5(128, 16, 64)
|
||||
self.stage2d = RSU6(128, 16, 64)
|
||||
self.stage1d = RSU7(128, 16, 64)
|
||||
|
||||
self.side1 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
self.side2 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
self.side3 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
self.side4 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
self.side5 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
self.side6 = nn.Conv2d(64, out_ch, 3, padding=1)
|
||||
|
||||
self.outconv = nn.Conv2d(6 * out_ch, out_ch, 1)
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
hx = x
|
||||
|
||||
# stage 1
|
||||
hx1 = self.stage1(hx)
|
||||
hx = self.pool12(hx1)
|
||||
|
||||
# stage 2
|
||||
hx2 = self.stage2(hx)
|
||||
hx = self.pool23(hx2)
|
||||
|
||||
# stage 3
|
||||
hx3 = self.stage3(hx)
|
||||
hx = self.pool34(hx3)
|
||||
|
||||
# stage 4
|
||||
hx4 = self.stage4(hx)
|
||||
hx = self.pool45(hx4)
|
||||
|
||||
# stage 5
|
||||
hx5 = self.stage5(hx)
|
||||
hx = self.pool56(hx5)
|
||||
|
||||
# stage 6
|
||||
hx6 = self.stage6(hx)
|
||||
hx6up = _upsample_like(hx6, hx5)
|
||||
|
||||
# decoder
|
||||
hx5d = self.stage5d(torch.cat((hx6up, hx5), 1))
|
||||
hx5dup = _upsample_like(hx5d, hx4)
|
||||
|
||||
hx4d = self.stage4d(torch.cat((hx5dup, hx4), 1))
|
||||
hx4dup = _upsample_like(hx4d, hx3)
|
||||
|
||||
hx3d = self.stage3d(torch.cat((hx4dup, hx3), 1))
|
||||
hx3dup = _upsample_like(hx3d, hx2)
|
||||
|
||||
hx2d = self.stage2d(torch.cat((hx3dup, hx2), 1))
|
||||
hx2dup = _upsample_like(hx2d, hx1)
|
||||
|
||||
hx1d = self.stage1d(torch.cat((hx2dup, hx1), 1))
|
||||
|
||||
# side output
|
||||
d1 = self.side1(hx1d)
|
||||
|
||||
d2 = self.side2(hx2d)
|
||||
d2 = _upsample_like(d2, d1)
|
||||
|
||||
d3 = self.side3(hx3d)
|
||||
d3 = _upsample_like(d3, d1)
|
||||
|
||||
d4 = self.side4(hx4d)
|
||||
d4 = _upsample_like(d4, d1)
|
||||
|
||||
d5 = self.side5(hx5d)
|
||||
d5 = _upsample_like(d5, d1)
|
||||
|
||||
d6 = self.side6(hx6)
|
||||
d6 = _upsample_like(d6, d1)
|
||||
|
||||
d0 = self.outconv(torch.cat((d1, d2, d3, d4, d5, d6), 1))
|
||||
|
||||
|
||||
return d0, d1, d2, d3, d4, d5, d6
|
||||
@@ -0,0 +1,12 @@
|
||||
import os.path as osp
|
||||
import os
|
||||
|
||||
|
||||
class parser(object):
|
||||
def __init__(self):
|
||||
|
||||
self.output = "./output" # output image folder path
|
||||
self.logs_dir = './logs'
|
||||
self.device = 'cuda:0'
|
||||
|
||||
opt = parser()
|
||||
@@ -0,0 +1,260 @@
|
||||
from network import U2NET
|
||||
|
||||
import os
|
||||
from PIL import Image
|
||||
import cv2
|
||||
import gdown
|
||||
import argparse
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import torchvision.transforms as transforms
|
||||
|
||||
from collections import OrderedDict
|
||||
from options import opt
|
||||
|
||||
import einops
|
||||
|
||||
def do_recolor(vis_seg_probs, n_classes):
|
||||
val = int(255 / n_classes)
|
||||
not_visible = (vis_seg_probs == 0).astype(dtype=np.uint8)
|
||||
not_visible = 1 - not_visible
|
||||
not_visible *= 255
|
||||
vis_seg_probs *= val
|
||||
ret = np.array((vis_seg_probs, not_visible, not_visible), np.uint8)
|
||||
ret = einops.rearrange(ret, 'c h w -> h w c')
|
||||
ret = cv2.cvtColor(ret, cv2.COLOR_HSV2RGB_FULL)
|
||||
return ret
|
||||
|
||||
|
||||
def load_checkpoint(model, checkpoint_path):
|
||||
if not os.path.exists(checkpoint_path):
|
||||
print("----No checkpoints at given path----")
|
||||
return
|
||||
model_state_dict = torch.load(checkpoint_path, map_location=torch.device("cpu"))
|
||||
new_state_dict = OrderedDict()
|
||||
for k, v in model_state_dict.items():
|
||||
name = k[7:] # remove `module.`
|
||||
new_state_dict[name] = v
|
||||
|
||||
model.load_state_dict(new_state_dict)
|
||||
print("----checkpoints loaded from path: {}----".format(checkpoint_path))
|
||||
return model
|
||||
|
||||
|
||||
def get_palette(num_cls):
|
||||
""" Returns the color map for visualizing the segmentation mask.
|
||||
Args:
|
||||
num_cls: Number of classes
|
||||
Returns:
|
||||
The color map
|
||||
"""
|
||||
n = num_cls
|
||||
palette = [0] * (n * 3)
|
||||
for j in range(0, n):
|
||||
lab = j
|
||||
palette[j * 3 + 0] = 0
|
||||
palette[j * 3 + 1] = 0
|
||||
palette[j * 3 + 2] = 0
|
||||
i = 0
|
||||
while lab:
|
||||
palette[j * 3 + 0] |= (((lab >> 0) & 1) << (7 - i))
|
||||
palette[j * 3 + 1] |= (((lab >> 1) & 1) << (7 - i))
|
||||
palette[j * 3 + 2] |= (((lab >> 2) & 1) << (7 - i))
|
||||
i += 1
|
||||
lab >>= 3
|
||||
return palette
|
||||
|
||||
|
||||
class Normalize_image(object):
|
||||
"""Normalize given tensor into given mean and standard dev
|
||||
|
||||
Args:
|
||||
mean (float): Desired mean to substract from tensors
|
||||
std (float): Desired std to divide from tensors
|
||||
"""
|
||||
|
||||
def __init__(self, mean, std):
|
||||
assert isinstance(mean, (float))
|
||||
if isinstance(mean, float):
|
||||
self.mean = mean
|
||||
|
||||
if isinstance(std, float):
|
||||
self.std = std
|
||||
|
||||
self.normalize_1 = transforms.Normalize(self.mean, self.std)
|
||||
self.normalize_3 = transforms.Normalize([self.mean] * 3, [self.std] * 3)
|
||||
self.normalize_18 = transforms.Normalize([self.mean] * 18, [self.std] * 18)
|
||||
|
||||
def __call__(self, image_tensor):
|
||||
if image_tensor.shape[0] == 1:
|
||||
return self.normalize_1(image_tensor)
|
||||
|
||||
elif image_tensor.shape[0] == 3:
|
||||
return self.normalize_3(image_tensor)
|
||||
|
||||
elif image_tensor.shape[0] == 18:
|
||||
return self.normalize_18(image_tensor)
|
||||
|
||||
else:
|
||||
assert "Please set proper channels! Normlization implemented only for 1, 3 and 18"
|
||||
|
||||
|
||||
|
||||
|
||||
def apply_transform(img):
|
||||
transforms_list = []
|
||||
transforms_list += [transforms.ToTensor()]
|
||||
transforms_list += [Normalize_image(0.5, 0.5)]
|
||||
transform_rgb = transforms.Compose(transforms_list)
|
||||
return transform_rgb(img)
|
||||
|
||||
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def generate_mask(input_image, net, device='cpu'):
|
||||
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')
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
print('#### DEBUG START ####')
|
||||
with torch.no_grad():
|
||||
output_tensor = net(image_tensor.to(device))
|
||||
print(output_tensor[0].shape)
|
||||
output_tensor = F.log_softmax(output_tensor[0], dim=1)
|
||||
output_tensor = torch.max(output_tensor, dim=1, keepdim=True)[1]
|
||||
output_tensor = torch.squeeze(output_tensor, dim=0)
|
||||
output_arr = output_tensor.cpu().numpy()
|
||||
|
||||
print(output_arr.shape)
|
||||
image_tmp = do_recolor(vis_seg_probs = output_arr.squeeze(0), n_classes = 4)
|
||||
print(image_tmp.shape)
|
||||
print('#### DEBUG STOP ####')
|
||||
|
||||
garment_path = os.path.join(output_dir, 'extracted_garment.png')
|
||||
cv2.imwrite(garment_path, cv2.cvtColor(src = image_tmp, code = cv2.COLOR_RGB2BGR))
|
||||
return image_tmp
|
||||
|
||||
# # 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
|
||||
# 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
|
||||
# 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')
|
||||
# extracted_garment.save(garment_path, format="PNG")
|
||||
|
||||
# return extracted_garment
|
||||
|
||||
|
||||
# def generate_mask(input_image, net, device='cpu'):
|
||||
# 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')
|
||||
# os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# with torch.no_grad():
|
||||
# output_tensor = net(image_tensor.to(device))
|
||||
# output_tensor = F.log_softmax(output_tensor[0], dim=1)
|
||||
# output_tensor = torch.max(output_tensor, dim=1, keepdim=True)[1]
|
||||
# output_tensor = torch.squeeze(output_tensor, dim=0)
|
||||
# output_arr = output_tensor.cpu().numpy()
|
||||
|
||||
# # 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
|
||||
|
||||
# # Convert binary mask to a 3-channel image to use as a mask
|
||||
|
||||
# # Ensure binary_mask is 2D
|
||||
# if binary_mask.ndim > 2:
|
||||
# binary_mask = binary_mask.squeeze() # Removes single-dimensional entries from the shape
|
||||
# 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)
|
||||
# binary_mask_3ch = binary_mask_img.convert('RGB') # Convert to RGB
|
||||
|
||||
# # Apply mask to the original image
|
||||
# original_img = img.resize(img_size) # Resize the processed image back to original size
|
||||
# extracted_garment = Image.new("RGB", original_img.size)
|
||||
# extracted_garment.paste(original_img, mask=binary_mask_img)
|
||||
|
||||
# # Save the garment image
|
||||
# garment_path = os.path.join(output_dir, 'extracted_garment.png')
|
||||
# extracted_garment.save(garment_path)
|
||||
|
||||
# return extracted_garment
|
||||
|
||||
|
||||
def check_or_download_model(file_path):
|
||||
if not os.path.exists(file_path):
|
||||
os.makedirs(os.path.dirname(file_path), exist_ok=True)
|
||||
url = "https://drive.google.com/uc?export=download&id=1qVv720hAd11JSCuIVJuqfjCGolwb1H8o"
|
||||
gdown.download(url, file_path, quiet=False)
|
||||
print("Model downloaded successfully.")
|
||||
else:
|
||||
print("Model already exists.")
|
||||
|
||||
|
||||
|
||||
def load_seg_model(checkpoint_path, device='cpu'):
|
||||
net = U2NET(in_ch=3, out_ch=4)
|
||||
check_or_download_model(checkpoint_path)
|
||||
net = load_checkpoint(net, checkpoint_path)
|
||||
net = net.to(device)
|
||||
net = net.eval()
|
||||
|
||||
return net
|
||||
|
||||
|
||||
def main(args):
|
||||
|
||||
device = 'cuda:0' if args.cuda else 'cpu'
|
||||
|
||||
# Create an instance of your model
|
||||
model = load_seg_model(args.checkpoint_path, device=device)
|
||||
|
||||
palette = get_palette(4)
|
||||
|
||||
img = Image.open(args.image).convert('RGB')
|
||||
|
||||
cloth_seg = generate_mask(img, net=model, palette=palette, device=device)
|
||||
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser(description='Help to set arguments for Cloth Segmentation.')
|
||||
parser.add_argument('--image', type=str, help='Path to the input image')
|
||||
parser.add_argument('--cuda', action='store_true', help='Enable CUDA (default: False)')
|
||||
parser.add_argument('--checkpoint_path', type=str, default='model/cloth_segm.pth', help='Path to the checkpoint file')
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
@@ -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"
|
||||
}
|
||||
@@ -0,0 +1,202 @@
|
||||
import numpy as np
|
||||
import cv2
|
||||
import math
|
||||
import torch
|
||||
|
||||
|
||||
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 smooth_step_plain(x):
|
||||
|
||||
if x < -1:
|
||||
return -1
|
||||
elif x <= 1:
|
||||
return math.sin(x * np.pi / 2.0)
|
||||
else:
|
||||
return 1
|
||||
|
||||
|
||||
def smooth_step_np(x):
|
||||
truths = np.logical_and(-1 < x, x < 1).astype(np.float32)
|
||||
x1 = np.clip(x, -1, 1)
|
||||
x2 = np.sin(x * np.pi / 2.0)
|
||||
ret = (truths * x2) + ((1 - truths) * x1)
|
||||
return ret
|
||||
|
||||
|
||||
def smooth_step_stretch(x, a, b):
|
||||
|
||||
if b < a:
|
||||
tmp = b
|
||||
b = a
|
||||
a = tmp
|
||||
|
||||
if a < 0:
|
||||
a = 0
|
||||
|
||||
if b > 1:
|
||||
b = 1
|
||||
|
||||
if a == b:
|
||||
a = 0
|
||||
b = 1
|
||||
|
||||
return smooth_step_np((2 * (x - a) / (b - a)) - 1)
|
||||
|
||||
|
||||
def get_light_layer(image,
|
||||
ref_r=255,
|
||||
ref_g=255,
|
||||
ref_b=255,
|
||||
do_scale=True,
|
||||
scale_a=0.0,
|
||||
scale_b=1.0):
|
||||
|
||||
sqmax = 3 * 255 * 255
|
||||
scalemax = math.sqrt(sqmax)
|
||||
|
||||
b = image[:, :, 0].astype(dtype=np.float32)
|
||||
g = image[:, :, 1].astype(dtype=np.float32)
|
||||
r = image[:, :, 2].astype(dtype=np.float32)
|
||||
|
||||
b2 = b * b
|
||||
g2 = g * g
|
||||
r2 = r * r
|
||||
|
||||
d2 = np.zeros(b2.shape, dtype=np.float32)
|
||||
d2 += sqmax - b2 - g2 - r2
|
||||
d = np.sqrt(d2)
|
||||
|
||||
ref_r2 = ref_r * ref_r
|
||||
ref_g2 = ref_g * ref_g
|
||||
ref_b2 = ref_b * ref_b
|
||||
ref_d2 = sqmax - ref_r2 - ref_g2 - ref_b2
|
||||
|
||||
ref_d = math.sqrt(ref_d2)
|
||||
|
||||
dot = (b * ref_b) + (g * ref_g) + (r * ref_r) + (d * ref_d)
|
||||
dot /= sqmax
|
||||
|
||||
if do_scale:
|
||||
dot = smooth_step_stretch(x=dot, a=scale_a, b=scale_b)
|
||||
|
||||
dot *= 255
|
||||
dot = dot.astype(np.uint8)
|
||||
|
||||
return dot
|
||||
|
||||
|
||||
class main_light_layer():
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", ),
|
||||
"ref_r": ("INT", {
|
||||
"default": 255,
|
||||
"min": 0,
|
||||
"max": 255,
|
||||
"step": 1
|
||||
}),
|
||||
"ref_g": ("INT", {
|
||||
"default": 255,
|
||||
"min": 0,
|
||||
"max": 255,
|
||||
"step": 1
|
||||
}),
|
||||
"ref_b": ("INT", {
|
||||
"default": 255,
|
||||
"min": 0,
|
||||
"max": 255,
|
||||
"step": 1
|
||||
}),
|
||||
"do_scale": (["enable", "disable"], ),
|
||||
"thresh_low": ("FLOAT", {
|
||||
"default": 0.6,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.01
|
||||
}),
|
||||
"thresh_high": ("FLOAT", {
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.01
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
FUNCTION = "run"
|
||||
RETURN_TYPES = ("MASK", )
|
||||
CATEGORY = "HackNode"
|
||||
|
||||
def run(
|
||||
self,
|
||||
image,
|
||||
ref_r,
|
||||
ref_g,
|
||||
ref_b,
|
||||
do_scale,
|
||||
thresh_low,
|
||||
thresh_high,
|
||||
):
|
||||
|
||||
do_scale = (do_scale == "enable")
|
||||
print('do_scale', do_scale)
|
||||
|
||||
image = from_torch_image(image)
|
||||
print('image.shape', image.shape)
|
||||
|
||||
batch_size = image.shape[0]
|
||||
print('batch_size', batch_size)
|
||||
|
||||
mask = []
|
||||
|
||||
for i in range(batch_size):
|
||||
|
||||
tmp_img = image[i]
|
||||
print('tmp_img.shape', tmp_img.shape)
|
||||
|
||||
tmp_mask = get_light_layer(
|
||||
tmp_img,
|
||||
ref_b,
|
||||
ref_g,
|
||||
ref_r,
|
||||
do_scale,
|
||||
scale_a=thresh_low,
|
||||
scale_b=thresh_high,
|
||||
)
|
||||
print('tmp_mask.shape', tmp_mask.shape)
|
||||
|
||||
mask.append(tmp_mask)
|
||||
|
||||
mask = np.array(mask)
|
||||
|
||||
mask = to_torch_image(mask)
|
||||
print(mask.shape)
|
||||
|
||||
return (mask, )
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
'main_light_layer': main_light_layer,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
'main_light_layer': 'main_light_layer',
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
import http.client
|
||||
import mimetypes
|
||||
import os
|
||||
import uuid
|
||||
import requests
|
||||
import numpy as np
|
||||
import torch
|
||||
import cv2
|
||||
from PIL import Image
|
||||
import io
|
||||
|
||||
class TRI3D_photoroom_bgremove_api:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE", ),
|
||||
},
|
||||
}
|
||||
|
||||
FUNCTION = "run"
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def run(self, images):
|
||||
import http.client
|
||||
import mimetypes
|
||||
import os
|
||||
import uuid
|
||||
import dotenv
|
||||
|
||||
dotenv.load_dotenv()
|
||||
|
||||
# Read the API key from the environment variable
|
||||
PHOTOROOM_API_KEY = os.getenv('PHOTOROOM_API_KEY','EMPTY')
|
||||
if PHOTOROOM_API_KEY == 'EMPTY':
|
||||
return images
|
||||
def tensor_to_cv2_img(tensor, remove_alpha=False):
|
||||
i = 255. * tensor.cpu().numpy() # This will give us (H, W, C)
|
||||
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
|
||||
|
||||
|
||||
# Please replace with your own apiKey
|
||||
|
||||
def remove_background(input_image_path, output_image_path,apiKey):
|
||||
# Define multipart boundary
|
||||
boundary = '----------{}'.format(uuid.uuid4().hex)
|
||||
|
||||
# Get mimetype of image
|
||||
content_type, _ = mimetypes.guess_type(input_image_path)
|
||||
if content_type is None:
|
||||
content_type = 'application/octet-stream' # Default type if guessing fails
|
||||
|
||||
# Prepare the POST data
|
||||
with open(input_image_path, 'rb') as f:
|
||||
image_data = f.read()
|
||||
filename = os.path.basename(input_image_path)
|
||||
|
||||
body = (
|
||||
f"--{boundary}\r\n"
|
||||
f"Content-Disposition: form-data; name=\"image_file\"; filename=\"{filename}\"\r\n"
|
||||
f"Content-Type: {content_type}\r\n\r\n"
|
||||
).encode('utf-8') + image_data + f"\r\n--{boundary}--\r\n".encode('utf-8')
|
||||
|
||||
# Set up the HTTP connection and headers
|
||||
conn = http.client.HTTPSConnection('sdk.photoroom.com')
|
||||
|
||||
headers = {
|
||||
'Content-Type': f'multipart/form-data; boundary={boundary}',
|
||||
'x-api-key': apiKey
|
||||
}
|
||||
|
||||
# Make the POST request
|
||||
conn.request('POST', '/v1/segment', body=body, headers=headers)
|
||||
response = conn.getresponse()
|
||||
|
||||
# Handle the response
|
||||
if response.status == 200:
|
||||
response_data = response.read()
|
||||
with open(output_image_path, 'wb') as out_f:
|
||||
out_f.write(response_data)
|
||||
print("Image saved to", output_image_path)
|
||||
else:
|
||||
print(f"Error: {response.status} - {response.reason}")
|
||||
print(response.read())
|
||||
|
||||
# Close the connection
|
||||
conn.close()
|
||||
|
||||
|
||||
|
||||
OUTPUT_FOLDER = "output/"
|
||||
batch_results = []
|
||||
for i in range(images.shape[0]):
|
||||
image = images[i]
|
||||
cv2_image = tensor_to_cv2_img(image)
|
||||
cv2_image = cv2.cvtColor(cv2_image, cv2.COLOR_BGR2RGB)
|
||||
import random
|
||||
random_number = random.randint(0, 100000)
|
||||
output_path = OUTPUT_FOLDER + f"output{i}_{random_number}.png"
|
||||
input_path = OUTPUT_FOLDER + f"input{i}_{random_number}.png"
|
||||
cv2.imwrite(input_path, cv2_image)
|
||||
remove_background(input_path, output_path, PHOTOROOM_API_KEY)
|
||||
|
||||
print(input_path, output_path)
|
||||
cv2_segm = cv2.imread(output_path, cv2.IMREAD_UNCHANGED)
|
||||
cv2_segm = cv2.cvtColor(cv2_segm, cv2.COLOR_BGRA2RGBA)
|
||||
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,)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -3,4 +3,9 @@ ninja
|
||||
pillow
|
||||
torch
|
||||
torchvision
|
||||
gdown
|
||||
transparent-background
|
||||
wget
|
||||
gdown
|
||||
matplotlib
|
||||
python-dotenv
|
||||
git+https://github.com/FacePerceiver/facer.git@main
|
||||
|
||||
|
After Width: | Height: | Size: 87 KiB |
|
After Width: | Height: | Size: 86 KiB |
|
After Width: | Height: | Size: 86 KiB |
|
After Width: | Height: | Size: 1.8 MiB |
|
After Width: | Height: | Size: 1.8 MiB |
|
After Width: | Height: | Size: 1.8 MiB |
|
After Width: | Height: | Size: 134 KiB |
@@ -0,0 +1,872 @@
|
||||
{
|
||||
"last_node_id": 28,
|
||||
"last_link_id": 35,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 2,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1968,
|
||||
-883
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
1
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"title": "Image1",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"image1.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1969,
|
||||
-514
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
2
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"title": "Image2",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"image2.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
1969,
|
||||
-151
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
3
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"title": "Image3",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"image3.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2370,
|
||||
-600
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
5
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"title": "ATR1",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"atr1.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2367,
|
||||
-236
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
6
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"title": "ATR2",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"atr2.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 14,
|
||||
"type": "ImageBatch",
|
||||
"pos": [
|
||||
2745,
|
||||
-584
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 46
|
||||
},
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image1",
|
||||
"type": "IMAGE",
|
||||
"link": 5
|
||||
},
|
||||
{
|
||||
"name": "image2",
|
||||
"type": "IMAGE",
|
||||
"link": 6
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
7
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ImageBatch"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 7,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2372,
|
||||
129
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
8
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"title": "ATR3",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"atr3.png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "ImageBatch",
|
||||
"pos": [
|
||||
2390,
|
||||
-848
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 46
|
||||
},
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image1",
|
||||
"type": "IMAGE",
|
||||
"link": 1
|
||||
},
|
||||
{
|
||||
"name": "image2",
|
||||
"type": "IMAGE",
|
||||
"link": 2
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
4
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ImageBatch"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 12,
|
||||
"type": "ImageBatch",
|
||||
"pos": [
|
||||
2981,
|
||||
-582
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 46
|
||||
},
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image1",
|
||||
"type": "IMAGE",
|
||||
"link": 7
|
||||
},
|
||||
{
|
||||
"name": "image2",
|
||||
"type": "IMAGE",
|
||||
"link": 8
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
10
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ImageBatch"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
2786,
|
||||
-268
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 314
|
||||
},
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
11
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"title": "Mask",
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"mask (30).png",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 15,
|
||||
"type": "RepeatImageBatch",
|
||||
"pos": [
|
||||
2793,
|
||||
-367
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 58
|
||||
},
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 11
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
18
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "RepeatImageBatch"
|
||||
},
|
||||
"widgets_values": [
|
||||
3
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 13,
|
||||
"type": "ImageBatch",
|
||||
"pos": [
|
||||
2636,
|
||||
-847
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 46
|
||||
},
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image1",
|
||||
"type": "IMAGE",
|
||||
"link": 4
|
||||
},
|
||||
{
|
||||
"name": "image2",
|
||||
"type": "IMAGE",
|
||||
"link": 3
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
16
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ImageBatch"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 26,
|
||||
"type": "Image To Mask",
|
||||
"pos": [
|
||||
3667,
|
||||
-11
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 58
|
||||
},
|
||||
"flags": {},
|
||||
"order": 15,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 32
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": [
|
||||
33
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "Image To Mask"
|
||||
},
|
||||
"widgets_values": [
|
||||
"intensity"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 27,
|
||||
"type": "InpaintPreprocessor",
|
||||
"pos": [
|
||||
3678,
|
||||
99
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 46
|
||||
},
|
||||
"flags": {},
|
||||
"order": 16,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 34
|
||||
},
|
||||
{
|
||||
"name": "mask",
|
||||
"type": "MASK",
|
||||
"link": 33
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
35
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "InpaintPreprocessor"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 17,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
3679,
|
||||
-294
|
||||
],
|
||||
"size": {
|
||||
"0": 686.6637573242188,
|
||||
"1": 246
|
||||
},
|
||||
"flags": {},
|
||||
"order": 14,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 19
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 28,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
3675,
|
||||
196
|
||||
],
|
||||
"size": {
|
||||
"0": 786.8681640625,
|
||||
"1": 246
|
||||
},
|
||||
"flags": {},
|
||||
"order": 17,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 35
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 16,
|
||||
"type": "PreviewImage",
|
||||
"pos": [
|
||||
3661,
|
||||
-574
|
||||
],
|
||||
"size": {
|
||||
"0": 714.0955810546875,
|
||||
"1": 246
|
||||
},
|
||||
"flags": {},
|
||||
"order": 13,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 17
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "PreviewImage"
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "tri3d-extract-parts-batch",
|
||||
"pos": [
|
||||
3280,
|
||||
-482
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 530
|
||||
},
|
||||
"flags": {},
|
||||
"order": 12,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "batch_images",
|
||||
"type": "IMAGE",
|
||||
"link": 16
|
||||
},
|
||||
{
|
||||
"name": "batch_segs",
|
||||
"type": "IMAGE",
|
||||
"link": 10
|
||||
},
|
||||
{
|
||||
"name": "batch_secondaries",
|
||||
"type": "IMAGE",
|
||||
"link": 18
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
17,
|
||||
34
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
19,
|
||||
32
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 1
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "tri3d-extract-parts-batch"
|
||||
},
|
||||
"widgets_values": [
|
||||
20,
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
true,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
1,
|
||||
2,
|
||||
0,
|
||||
11,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
2,
|
||||
3,
|
||||
0,
|
||||
11,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
3,
|
||||
4,
|
||||
0,
|
||||
13,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
4,
|
||||
11,
|
||||
0,
|
||||
13,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
5,
|
||||
5,
|
||||
0,
|
||||
14,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
6,
|
||||
6,
|
||||
0,
|
||||
14,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
7,
|
||||
14,
|
||||
0,
|
||||
12,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
8,
|
||||
7,
|
||||
0,
|
||||
12,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
10,
|
||||
12,
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
11,
|
||||
8,
|
||||
0,
|
||||
15,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
16,
|
||||
13,
|
||||
0,
|
||||
1,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
17,
|
||||
1,
|
||||
0,
|
||||
16,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
18,
|
||||
15,
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
19,
|
||||
1,
|
||||
1,
|
||||
17,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
32,
|
||||
1,
|
||||
1,
|
||||
26,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
33,
|
||||
26,
|
||||
0,
|
||||
27,
|
||||
1,
|
||||
"MASK"
|
||||
],
|
||||
[
|
||||
34,
|
||||
1,
|
||||
0,
|
||||
27,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
35,
|
||||
27,
|
||||
0,
|
||||
28,
|
||||
0,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
#!/usr/bin/python3
|
||||
import torch
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
#!/usr/bin/python3
|
||||
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 scaled_paste(
|
||||
image_background,
|
||||
image_foreground,
|
||||
mask_foreground,
|
||||
scale_factor,
|
||||
height_factor=1.2,
|
||||
):
|
||||
|
||||
print('DEBUG scaled_paste 0 ', image_background.shape,
|
||||
image_foreground.shape, mask_foreground.shape, scale_factor,
|
||||
height_factor)
|
||||
|
||||
height = image_foreground.shape[0] * height_factor
|
||||
|
||||
print('DEBUG scaled_paste 1 ', height)
|
||||
|
||||
max_0 = max(image_background.shape[0], height)
|
||||
max_1 = max(image_background.shape[1], image_foreground.shape[1])
|
||||
|
||||
print('DEBUG scaled_paste 2 ', max_0, max_1)
|
||||
|
||||
ratio_0 = max_0 / image_background.shape[0]
|
||||
ratio_1 = max_1 / image_background.shape[1]
|
||||
ratio_max = max(ratio_0, ratio_1) * scale_factor
|
||||
|
||||
print('DEBUG scaled_paste 2 ', ratio_0, ratio_1, ratio_max)
|
||||
|
||||
size_0 = int(image_background.shape[0] * ratio_max) + 1
|
||||
size_1 = int(image_background.shape[1] * ratio_max) + 1
|
||||
|
||||
print('DEBUG scaled_paste 3 ', size_0, size_1)
|
||||
|
||||
image_background = cv2.resize(image_background, (size_1, size_0),
|
||||
cv2.INTER_CUBIC)
|
||||
|
||||
print('DEBUG scaled_paste 4 ', image_background.shape)
|
||||
|
||||
end_0 = int(image_background.shape[0])
|
||||
begin_0 = int(end_0 - height)
|
||||
end_0 = int(begin_0 + image_foreground.shape[0])
|
||||
|
||||
print('DEBUG scaled_paste 5 ', begin_0, end_0)
|
||||
|
||||
end_1 = image_background.shape[1]
|
||||
begin_1 = end_1 - image_foreground.shape[1]
|
||||
begin_1 = int(begin_1 / 2)
|
||||
end_1 = int(begin_1 + image_foreground.shape[1])
|
||||
|
||||
print('DEBUG scaled_paste 6 ', begin_1, end_1)
|
||||
|
||||
image_reference = image_background[begin_0:end_0, begin_1:end_1, :]
|
||||
|
||||
for i in range(3):
|
||||
image_reference[:, :,
|
||||
i] = (mask_foreground * image_foreground[:, :, i]) + (
|
||||
(1 - mask_foreground) * image_reference[:, :, i])
|
||||
|
||||
return image_background
|
||||
|
||||
|
||||
#!/usr/bin/python3
|
||||
class main_scaled_paste():
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image_background": ("IMAGE", ),
|
||||
"image_foreground": ("IMAGE", ),
|
||||
"mask_foreground": ("MASK", ),
|
||||
"scale_factor": ("FLOAT", {
|
||||
"default": 1.2,
|
||||
"min": 1,
|
||||
"max": 10,
|
||||
"step": 0.05
|
||||
}),
|
||||
"height_factor": ("FLOAT", {
|
||||
"default": 1.01,
|
||||
"min": 1,
|
||||
"max": 8,
|
||||
"step": 0.05
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
FUNCTION = "run"
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def run(
|
||||
self,
|
||||
image_background,
|
||||
image_foreground,
|
||||
mask_foreground,
|
||||
scale_factor,
|
||||
height_factor,
|
||||
):
|
||||
|
||||
print('DEBUG 0 ', image_background.shape, image_foreground.shape,
|
||||
mask_foreground.shape)
|
||||
|
||||
image_background = from_torch_image(image_background)
|
||||
image_foreground = from_torch_image(image_foreground)
|
||||
mask_foreground = mask_foreground.cpu().numpy()
|
||||
|
||||
image_output = scaled_paste(
|
||||
image_background[0],
|
||||
image_foreground[0],
|
||||
mask_foreground[0],
|
||||
scale_factor,
|
||||
height_factor,
|
||||
)
|
||||
|
||||
print('DEBUG 1 ', image_output.shape)
|
||||
|
||||
image_output = to_torch_image(image=image_output)
|
||||
|
||||
print('DEBUG 2 ', image_output.shape)
|
||||
|
||||
image_output = image_output.unsqueeze(0)
|
||||
|
||||
print('DEBUG 3 ', image_output.shape)
|
||||
|
||||
return (image_output, )
|
||||
|
||||
|
||||
#!/usr/bin/python3
|
||||
# mask = cv2.imread('/home/asd/DATASETS/BG_SWAP_HACK_TEST/FOREGROUND_MASK.png',
|
||||
# cv2.IMREAD_GRAYSCALE)
|
||||
|
||||
# mask = mask.astype(dtype=np.float32) / 255.0
|
||||
|
||||
# image_background = scaled_paste(
|
||||
# image_background=cv2.imread(
|
||||
# '/home/asd/DATASETS/BG_SWAP_HACK_TEST/BACKGROUND_DEPTH.png',
|
||||
# cv2.IMREAD_COLOR),
|
||||
# image_foreground=cv2.imread(
|
||||
# '/home/asd/DATASETS/BG_SWAP_HACK_TEST/FOREGROUND_DEPTH.png',
|
||||
# cv2.IMREAD_COLOR),
|
||||
# mask_foreground=mask,
|
||||
# scale_factor=2,
|
||||
# height_factor=1.05,
|
||||
# )
|
||||
|
||||
# cv2.imwrite('tmp.png', image_background)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
'main_scaled_paste': main_scaled_paste,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
'main_scaled_paste': 'main_scaled_paste',
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
#!/usr/bin/python3
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import einops
|
||||
|
||||
import facer
|
||||
|
||||
|
||||
def load_image(image_path):
|
||||
image = cv2.imread(image_path, cv2.IMREAD_COLOR)
|
||||
image = cv2.cvtColor(image, code=cv2.COLOR_BGR2RGB)
|
||||
image = torch.from_numpy(image).to(dtype=torch.float32) / 255.0
|
||||
return image
|
||||
|
||||
|
||||
def do_recolor(vis_seg_probs, n_classes):
|
||||
val = int(255 / n_classes)
|
||||
vis_seg_probs = vis_seg_probs.cpu().detach().numpy()
|
||||
not_visible = (vis_seg_probs == 0).astype(dtype=np.uint8)
|
||||
not_visible = 1 - not_visible
|
||||
not_visible *= 255
|
||||
vis_seg_probs *= val
|
||||
ret = np.array((vis_seg_probs, not_visible, not_visible), np.uint8)
|
||||
ret = einops.rearrange(ret, 'c h w -> h w c')
|
||||
ret = cv2.cvtColor(ret, cv2.COLOR_HSV2BGR_FULL)
|
||||
return ret
|
||||
|
||||
|
||||
def detect_face_from_tensor(image):
|
||||
device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
image *= 255
|
||||
image = image.to(dtype=torch.uint8)
|
||||
image = facer.hwc2bchw(image).to(device=device)
|
||||
face_detector = facer.face_detector('retinaface/mobilenet', device=device)
|
||||
|
||||
with torch.inference_mode():
|
||||
faces = face_detector(image)
|
||||
|
||||
face_parser = facer.face_parser(
|
||||
'farl/lapa/448', device=device) # optional "farl/celebm/448"
|
||||
|
||||
with torch.inference_mode():
|
||||
faces = face_parser(image, faces)
|
||||
|
||||
seg_logits = faces['seg']['logits']
|
||||
num_faces = seg_logits.shape[0]
|
||||
print(num_faces)
|
||||
|
||||
if num_faces >= 1:
|
||||
|
||||
seg_probs = seg_logits.softmax(dim=1) # nfaces x nclasses x h x w
|
||||
n_classes = seg_probs.size(1)
|
||||
|
||||
vis_seg_probs = seg_probs.argmax(dim=1)
|
||||
vis_seg_probs = einops.einsum(vis_seg_probs, 'b h w -> h w')
|
||||
|
||||
return (vis_seg_probs, n_classes, num_faces)
|
||||
|
||||
else:
|
||||
|
||||
vis_seg_probs = torch.zeros((image.shape[0], image.shape[1]),
|
||||
dtype=torch.int64)
|
||||
|
||||
n_classes = 11
|
||||
|
||||
return (vis_seg_probs, n_classes, num_faces)
|
||||
|
||||
|
||||
def full_work_wrapper(image):
|
||||
|
||||
try:
|
||||
res, n_classes, num_faces = detect_face_from_tensor(image)
|
||||
except:
|
||||
res = torch.zeros((image.shape[0], image.shape[1]), dtype=torch.int64)
|
||||
n_classes = 11
|
||||
tup = do_recolor(res, n_classes)
|
||||
return tup
|
||||
|
||||
|
||||
def run_slave(input_image_path, output_image_path, tmp_file_path):
|
||||
import os
|
||||
|
||||
EXEC_STRING = '''
|
||||
import os
|
||||
|
||||
try:
|
||||
del os.environ['AUX_ANNOTATOR_CKPTS_PATH']
|
||||
os.unsetenv('AUX_ANNOTATOR_CKPTS_PATH')
|
||||
except:
|
||||
print('Failed to unset AUX_ANNOTATOR_CKPTS_PATH')
|
||||
try:
|
||||
del os.environ['AUX_ORT_PROVIDERS']
|
||||
os.unsetenv('AUX_ORT_PROVIDERS')
|
||||
except:
|
||||
print('Failed to unset AUX_ORT_PROVIDERS')
|
||||
try:
|
||||
del os.environ['AUX_TEMP_DIR']
|
||||
os.unsetenv('AUX_TEMP_DIR')
|
||||
except:
|
||||
print('Failed to unset AUX_TEMP_DIR')
|
||||
try:
|
||||
del os.environ['AUX_USE_SYMLINKS']
|
||||
os.unsetenv('AUX_USE_SYMLINKS')
|
||||
except:
|
||||
print('Failed to unset AUX_USE_SYMLINKS')
|
||||
try:
|
||||
del os.environ['CUBLAS_WORKSPACE_CONFIG']
|
||||
os.unsetenv('CUBLAS_WORKSPACE_CONFIG')
|
||||
except:
|
||||
print('Failed to unset CUBLAS_WORKSPACE_CONFIG')
|
||||
try:
|
||||
del os.environ['CUDA_MODULE_LOADING']
|
||||
os.unsetenv('CUDA_MODULE_LOADING')
|
||||
except:
|
||||
print('Failed to unset CUDA_MODULE_LOADING')
|
||||
try:
|
||||
del os.environ['DWPOSE_ONNXRT_CHECKED']
|
||||
os.unsetenv('DWPOSE_ONNXRT_CHECKED')
|
||||
except:
|
||||
print('Failed to unset DWPOSE_ONNXRT_CHECKED')
|
||||
try:
|
||||
del os.environ['KINETO_LOG_LEVEL']
|
||||
os.unsetenv('KINETO_LOG_LEVEL')
|
||||
except:
|
||||
print('Failed to unset KINETO_LOG_LEVEL')
|
||||
try:
|
||||
del os.environ['KMP_DUPLICATE_LIB_OK']
|
||||
os.unsetenv('KMP_DUPLICATE_LIB_OK')
|
||||
except:
|
||||
print('Failed to unset KMP_DUPLICATE_LIB_OK')
|
||||
try:
|
||||
del os.environ['KMP_INIT_AT_FORK']
|
||||
os.unsetenv('KMP_INIT_AT_FORK')
|
||||
except:
|
||||
print('Failed to unset KMP_INIT_AT_FORK')
|
||||
try:
|
||||
del os.environ['PYTORCH_CUDA_ALLOC_CONF']
|
||||
os.unsetenv('PYTORCH_CUDA_ALLOC_CONF')
|
||||
except:
|
||||
print('Failed to unset PYTORCH_CUDA_ALLOC_CONF')
|
||||
try:
|
||||
del os.environ['PYTORCH_ENABLE_MPS_FALLBACK']
|
||||
os.unsetenv('PYTORCH_ENABLE_MPS_FALLBACK')
|
||||
except:
|
||||
print('Failed to unset PYTORCH_ENABLE_MPS_FALLBACK')
|
||||
try:
|
||||
del os.environ['PYTORCH_NVML_BASED_CUDA_CHECK']
|
||||
os.unsetenv('PYTORCH_NVML_BASED_CUDA_CHECK')
|
||||
except:
|
||||
print('Failed to unset PYTORCH_NVML_BASED_CUDA_CHECK')
|
||||
try:
|
||||
del os.environ['TF_CPP_MIN_LOG_LEVEL']
|
||||
os.unsetenv('TF_CPP_MIN_LOG_LEVEL')
|
||||
except:
|
||||
print('Failed to unset TF_CPP_MIN_LOG_LEVEL')
|
||||
try:
|
||||
del os.environ['TOKENIZERS_PARALLELISM']
|
||||
os.unsetenv('TOKENIZERS_PARALLELISM')
|
||||
except:
|
||||
print('Failed to unset TOKENIZERS_PARALLELISM')
|
||||
try:
|
||||
del os.environ['TORCH_CPP_LOG_LEVEL']
|
||||
os.unsetenv('TORCH_CPP_LOG_LEVEL')
|
||||
except:
|
||||
print('Failed to unset TORCH_CPP_LOG_LEVEL')
|
||||
|
||||
import torch
|
||||
import facer
|
||||
import cv2
|
||||
import einops
|
||||
import numpy as np
|
||||
import sys
|
||||
|
||||
|
||||
def load_image(image_path):
|
||||
image = cv2.imread(image_path, cv2.IMREAD_COLOR)
|
||||
image = cv2.cvtColor(image, code=cv2.COLOR_BGR2RGB)
|
||||
image = torch.from_numpy(image).to(dtype=torch.float32) / 255.0
|
||||
return image
|
||||
|
||||
|
||||
def do_recolor(vis_seg_probs, n_classes):
|
||||
val = int(255 / n_classes)
|
||||
vis_seg_probs = vis_seg_probs.cpu().detach().numpy()
|
||||
not_visible = (vis_seg_probs == 0).astype(dtype=np.uint8)
|
||||
not_visible = 1 - not_visible
|
||||
not_visible *= 255
|
||||
vis_seg_probs *= val
|
||||
ret = np.array((vis_seg_probs, not_visible, not_visible), np.uint8)
|
||||
ret = einops.rearrange(ret, 'c h w -> h w c')
|
||||
ret = cv2.cvtColor(ret, cv2.COLOR_HSV2BGR_FULL)
|
||||
return ret
|
||||
|
||||
|
||||
def detect_face_from_tensor(image):
|
||||
device = 'cuda' if torch.cuda.is_available() else 'cpu'
|
||||
image *= 255
|
||||
image = image.to(dtype=torch.uint8)
|
||||
image = facer.hwc2bchw(image).to(device=device)
|
||||
face_detector = facer.face_detector('retinaface/mobilenet', device=device)
|
||||
|
||||
with torch.inference_mode():
|
||||
faces = face_detector(image)
|
||||
|
||||
face_parser = facer.face_parser(
|
||||
'farl/lapa/448', device=device) # optional "farl/celebm/448"
|
||||
|
||||
with torch.inference_mode():
|
||||
faces = face_parser(image, faces)
|
||||
|
||||
seg_logits = faces['seg']['logits']
|
||||
seg_probs = seg_logits.softmax(dim=1) # nfaces x nclasses x h x w
|
||||
n_classes = seg_probs.size(1)
|
||||
|
||||
vis_seg_probs = seg_probs.argmax(dim=1)
|
||||
vis_seg_probs = einops.einsum(vis_seg_probs, 'b h w -> h w')
|
||||
return (vis_seg_probs, n_classes)
|
||||
|
||||
|
||||
def full_work_wrapper(image):
|
||||
try:
|
||||
res, n_classes = detect_face_from_tensor(image)
|
||||
tup = do_recolor(res, n_classes)
|
||||
except:
|
||||
print('Warning: Failed to find a face.')
|
||||
tup = np.zeros(image.shape, dtype=np.uint8)
|
||||
return tup
|
||||
|
||||
tup = full_work_wrapper(image=load_image(image_path=sys.argv[1]))
|
||||
cv2.imwrite(sys.argv[2], tup)
|
||||
'''
|
||||
|
||||
with open(tmp_file_path, 'w', encoding='utf-8') as f:
|
||||
f.write(EXEC_STRING)
|
||||
|
||||
CMD = 'env > ~/env.txt ; python3 ' + tmp_file_path + ' ' + input_image_path + ' ' + output_image_path
|
||||
|
||||
print(CMD)
|
||||
os.system(CMD)
|
||||
|
||||
|
||||
def run_slave_tensor(image):
|
||||
|
||||
import tempfile
|
||||
import cv2
|
||||
import os
|
||||
|
||||
device = image.device
|
||||
outtype = image.dtype
|
||||
|
||||
path_dir = tempfile.TemporaryDirectory(
|
||||
suffix='.dir',
|
||||
prefix='facer.',
|
||||
dir=None,
|
||||
ignore_cleanup_errors=False,
|
||||
)
|
||||
|
||||
path_input = path_dir.name + '/input.png'
|
||||
path_output = path_dir.name + '/output.png'
|
||||
path_source = path_dir.name + '/exec.py'
|
||||
|
||||
image = image.detach().cpu().numpy() * 255.0
|
||||
image = image.astype(dtype=np.uint8)
|
||||
image = cv2.cvtColor(src=image, code=cv2.COLOR_RGB2BGR)
|
||||
cv2.imwrite(path_input, image)
|
||||
|
||||
run_slave(input_image_path=path_input,
|
||||
output_image_path=path_output,
|
||||
tmp_file_path=path_source)
|
||||
|
||||
os.unlink(path_input)
|
||||
os.unlink(path_source)
|
||||
image = cv2.imread(path_output, cv2.IMREAD_COLOR)
|
||||
os.unlink(path_output)
|
||||
os.rmdir(path_dir.name)
|
||||
# image = cv2.cvtColor(src=image, code=cv2.COLOR_BGR2RGB)
|
||||
image = image.astype(np.float32) / 255.0
|
||||
# image = torch.from_numpy(image).to(dtype=outtype, device=device) / 255.0
|
||||
return image
|
||||
|
||||
|
||||
class main_face_segment():
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", ),
|
||||
"to_run": ("BOOLEAN", )
|
||||
},
|
||||
}
|
||||
|
||||
FUNCTION = "run"
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def run(self, image, to_run):
|
||||
if to_run:
|
||||
batch_size = image.shape[0]
|
||||
ret = []
|
||||
for i in range(batch_size):
|
||||
ret.append(run_slave_tensor(image[i].clone()))
|
||||
# ret.append(full_work_wrapper(image[i].clone()))
|
||||
|
||||
ret = np.array(ret)
|
||||
|
||||
ret = torch.from_numpy(ret).to(dtype=image.dtype,
|
||||
device=image.device)
|
||||
|
||||
return (ret, )
|
||||
else:
|
||||
return (torch.zeros_like(image), )
|
||||
@@ -0,0 +1,610 @@
|
||||
#!/usr/bin/python3
|
||||
import cv2
|
||||
import math
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
#!/usr/bin/python3
|
||||
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 do_custom_threshhold(image, value):
|
||||
image = image.astype(dtype=np.float64)
|
||||
image = 255 * (image - value) / (255 - value)
|
||||
image = np.clip(image, 0, 255)
|
||||
image = image.astype(dtype=np.uint8)
|
||||
return image
|
||||
|
||||
|
||||
def scaled_paste(
|
||||
image_background,
|
||||
image_foreground,
|
||||
mask_foreground,
|
||||
scale_factor=1.2,
|
||||
height_factor=1.05,
|
||||
):
|
||||
|
||||
height = image_foreground.shape[0] * height_factor
|
||||
|
||||
max_0 = max(image_background.shape[0], height)
|
||||
max_1 = max(image_background.shape[1], image_foreground.shape[1])
|
||||
|
||||
ratio_0 = max_0 / image_background.shape[0]
|
||||
ratio_1 = max_1 / image_background.shape[1]
|
||||
ratio_max = max(ratio_0, ratio_1) * scale_factor
|
||||
|
||||
size_0 = int(image_background.shape[0] * ratio_max) + 1
|
||||
size_1 = int(image_background.shape[1] * ratio_max) + 1
|
||||
|
||||
image_background = cv2.resize(image_background, (size_1, size_0),
|
||||
cv2.INTER_CUBIC)
|
||||
|
||||
end_0 = int(image_background.shape[0])
|
||||
begin_0 = int(end_0 - height)
|
||||
end_0 = int(begin_0 + image_foreground.shape[0])
|
||||
|
||||
end_1 = image_background.shape[1]
|
||||
begin_1 = end_1 - image_foreground.shape[1]
|
||||
begin_1 = int(begin_1 / 2)
|
||||
end_1 = int(begin_1 + image_foreground.shape[1])
|
||||
|
||||
image_reference = image_background[begin_0:end_0, begin_1:end_1, :]
|
||||
|
||||
for i in range(3):
|
||||
image_reference[:, :,
|
||||
i] = (mask_foreground * image_foreground[:, :, i]) + (
|
||||
(1 - mask_foreground) * image_reference[:, :, i])
|
||||
|
||||
return image_background
|
||||
|
||||
|
||||
def do_bg_swap(
|
||||
bkg_image,
|
||||
subject_image,
|
||||
mask_image,
|
||||
threshhold_hist,
|
||||
scale_factor=1.2,
|
||||
height_factor=1.05,
|
||||
):
|
||||
|
||||
blank_subject_image = np.zeros(subject_image.shape, dtype=np.uint8)
|
||||
blank_subject_image += 255
|
||||
|
||||
blank_background_image = np.zeros(bkg_image.shape, dtype=np.uint8)
|
||||
|
||||
blank_subject_mask = np.zeros(
|
||||
(subject_image.shape[0], subject_image.shape[1]), dtype=np.float64)
|
||||
|
||||
blank_subject_mask += 1
|
||||
|
||||
mask_image_3channel = subject_image.copy()
|
||||
for i in range(3):
|
||||
mask_image_3channel[:, :, i] = mask_image
|
||||
|
||||
mask_image = mask_image.astype(dtype=np.float64)
|
||||
mask_image /= 255.0
|
||||
|
||||
result_image = scaled_paste(
|
||||
image_background=bkg_image,
|
||||
image_foreground=subject_image,
|
||||
mask_foreground=mask_image,
|
||||
scale_factor=scale_factor,
|
||||
height_factor=height_factor,
|
||||
)
|
||||
|
||||
luminosity_image = scaled_paste(
|
||||
image_background=blank_background_image + 255,
|
||||
image_foreground=subject_image,
|
||||
mask_foreground=blank_subject_mask,
|
||||
scale_factor=scale_factor,
|
||||
height_factor=height_factor,
|
||||
)
|
||||
|
||||
final_mask = scaled_paste(
|
||||
image_background=blank_background_image,
|
||||
image_foreground=mask_image_3channel,
|
||||
mask_foreground=mask_image,
|
||||
scale_factor=scale_factor,
|
||||
height_factor=height_factor,
|
||||
)
|
||||
|
||||
result_image_lab = cv2.cvtColor(src=result_image, code=cv2.COLOR_RGB2LAB)
|
||||
|
||||
luminosity_image_lab = cv2.cvtColor(src=luminosity_image,
|
||||
code=cv2.COLOR_RGB2LAB)[:, :, 0]
|
||||
|
||||
luminosity_image_lab_flip = 255 - luminosity_image_lab
|
||||
|
||||
luminosity_image_lab_flip = do_custom_threshhold(
|
||||
image=luminosity_image_lab_flip, value=threshhold_hist)
|
||||
|
||||
luminosity_image_lab_flip_full = luminosity_image_lab_flip.copy()
|
||||
|
||||
luminosity_image_lab_flip *= 1 - (final_mask[:, :, 0]
|
||||
> 127.5).astype(dtype=np.uint8)
|
||||
|
||||
for i in range(3):
|
||||
result_image[:, :,
|
||||
i] = (result_image[:, :, i] *
|
||||
(1 - (luminosity_image_lab_flip / 255.0))).astype(
|
||||
dtype=np.uint8)
|
||||
|
||||
return (result_image, luminosity_image_lab_flip_full)
|
||||
|
||||
|
||||
def find_threshold(image_input, threshold=0.0001):
|
||||
|
||||
image_input_L = cv2.cvtColor(image_input, cv2.COLOR_RGB2LAB)[:, :,
|
||||
0].flatten()
|
||||
image_input_L = 255 - image_input_L
|
||||
|
||||
hist = np.histogram(image_input_L, range(0, 256, 1))
|
||||
|
||||
values = hist[0]
|
||||
values = values.astype(dtype=np.float64)
|
||||
values /= len(image_input_L)
|
||||
|
||||
for i in range(0, values.shape[0], 1):
|
||||
|
||||
lhd = 0
|
||||
rhd = 0
|
||||
|
||||
if i > 0:
|
||||
lhd = values[i] - values[i - 1]
|
||||
|
||||
if i < values.shape[0] - 1:
|
||||
rhd = values[i + 1] - values[i]
|
||||
|
||||
print(lhd, rhd)
|
||||
|
||||
if max(lhd, rhd) > threshold:
|
||||
return i
|
||||
|
||||
|
||||
def get_mu_sigma(array_input, mask_input):
|
||||
|
||||
array_input = array_input.astype(dtype=np.float32).flatten()
|
||||
mask_input = mask_input.astype(dtype=np.float32).flatten()
|
||||
|
||||
sum = np.sum(mask_input)
|
||||
mean = np.sum(array_input * mask_input) / sum
|
||||
|
||||
array_input -= mean
|
||||
array_input *= mask_input
|
||||
sigma = math.sqrt(np.sum(np.square(array_input)) / sum)
|
||||
|
||||
return mean, sigma
|
||||
|
||||
|
||||
def renormalize_array_main(array_input, mask_input, mu, sigma):
|
||||
array_input_original = array_input.copy()
|
||||
in_mu, in_sigma = get_mu_sigma(array_input, mask_input)
|
||||
array_input = (((array_input - in_mu) / in_sigma) * sigma) + mu
|
||||
|
||||
array_input_original = (array_input_original *
|
||||
(1 - mask_input)) + (array_input * mask_input)
|
||||
|
||||
return array_input_original
|
||||
|
||||
|
||||
#!/usr/bin/python3
|
||||
class simple_bg_swap:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"bkg_image": ("IMAGE", ),
|
||||
"subject_image": ("IMAGE", ),
|
||||
"subject_mask": ("MASK", ),
|
||||
"threshhold_hist": (
|
||||
"INT",
|
||||
{
|
||||
"default": 150,
|
||||
"min": 0, #Minimum value
|
||||
"max": 255, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display":
|
||||
"number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"scale_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.2,
|
||||
"min": 0.0,
|
||||
"max": 10.0,
|
||||
"step": 0.01,
|
||||
"round":
|
||||
0.001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding.
|
||||
"display": "number"
|
||||
}),
|
||||
"height_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.05,
|
||||
"min": 1.0,
|
||||
"max": 8.0,
|
||||
"step": 0.01,
|
||||
"round":
|
||||
0.001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding.
|
||||
"display": "number"
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
"IMAGE",
|
||||
"MASK",
|
||||
)
|
||||
|
||||
RETURN_NAMES = (
|
||||
"output bg swapped image",
|
||||
"shadow layer",
|
||||
)
|
||||
|
||||
FUNCTION = "test"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def test(
|
||||
self,
|
||||
bkg_image,
|
||||
subject_image,
|
||||
subject_mask,
|
||||
threshhold_hist,
|
||||
scale_factor,
|
||||
height_factor,
|
||||
):
|
||||
|
||||
mask_image = subject_mask
|
||||
|
||||
bkg_image = from_torch_image(image=bkg_image)
|
||||
subject_image = from_torch_image(image=subject_image)
|
||||
mask_image = from_torch_image(image=mask_image)
|
||||
|
||||
batch_size = bkg_image.shape[0]
|
||||
|
||||
ret = []
|
||||
ret_lum = []
|
||||
|
||||
if (subject_image.shape[0] == batch_size) and (mask_image.shape[0]
|
||||
== batch_size):
|
||||
|
||||
for i in range(batch_size):
|
||||
|
||||
result, luminosity = do_bg_swap(
|
||||
bkg_image[i],
|
||||
subject_image[i],
|
||||
mask_image[i],
|
||||
threshhold_hist,
|
||||
scale_factor,
|
||||
height_factor,
|
||||
)
|
||||
|
||||
result = to_torch_image(result)
|
||||
result = result.unsqueeze(0)
|
||||
ret.append(result)
|
||||
|
||||
luminosity = to_torch_image(luminosity)
|
||||
luminosity = luminosity.unsqueeze(0)
|
||||
ret_lum.append(luminosity)
|
||||
|
||||
else:
|
||||
|
||||
print(
|
||||
'input format is not correct, got different batch sizes for each input image'
|
||||
)
|
||||
|
||||
ret = torch.cat(ret, dim=0)
|
||||
ret_lum = torch.cat(ret_lum, dim=0)
|
||||
|
||||
return (
|
||||
ret,
|
||||
ret_lum,
|
||||
)
|
||||
|
||||
|
||||
class get_threshold_for_bg_swap:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"subject_image": ("IMAGE", ),
|
||||
"gradient_threshold": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0001,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.00001,
|
||||
"round":
|
||||
0.000001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding.
|
||||
"display": "number"
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT", )
|
||||
RETURN_NAMES = ("output histogram threshold", )
|
||||
|
||||
FUNCTION = "test"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def test(
|
||||
self,
|
||||
subject_image,
|
||||
gradient_threshold,
|
||||
):
|
||||
|
||||
subject_image = from_torch_image(image=subject_image)
|
||||
return (find_threshold(subject_image[0],
|
||||
threshold=gradient_threshold), )
|
||||
|
||||
|
||||
class RGB_2_LAB:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_RGB_image": ("IMAGE", ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK", "MASK", "MASK")
|
||||
RETURN_NAMES = ("L", "A", "B")
|
||||
|
||||
FUNCTION = "test"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def test(self, input_RGB_image):
|
||||
print('input_RGB_image.shape', input_RGB_image.shape)
|
||||
input_RGB_image = from_torch_image(image=input_RGB_image)
|
||||
ret_L = []
|
||||
ret_A = []
|
||||
ret_B = []
|
||||
for i in range(input_RGB_image.shape[0]):
|
||||
tmp = cv2.cvtColor(input_RGB_image[i], cv2.COLOR_RGB2LAB)
|
||||
ret_L.append(to_torch_image(image=tmp[:, :, 0]).unsqueeze(0))
|
||||
ret_A.append(to_torch_image(image=tmp[:, :, 1]).unsqueeze(0))
|
||||
ret_B.append(to_torch_image(image=tmp[:, :, 2]).unsqueeze(0))
|
||||
|
||||
ret_L = torch.cat(ret_L, dim=0)
|
||||
ret_A = torch.cat(ret_A, dim=0)
|
||||
ret_B = torch.cat(ret_B, dim=0)
|
||||
|
||||
print(
|
||||
'LAB output',
|
||||
ret_L.shape,
|
||||
ret_A.shape,
|
||||
ret_B.shape,
|
||||
)
|
||||
|
||||
return (ret_L, ret_A, ret_B)
|
||||
|
||||
|
||||
class LAB_2_RGB:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_L": ("MASK", ),
|
||||
"input_A": ("MASK", ),
|
||||
"input_B": ("MASK", ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
RETURN_NAMES = ("Output RGB image", )
|
||||
|
||||
FUNCTION = "test"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def test(self, input_L, input_A, input_B):
|
||||
|
||||
batch_size = input_L.shape[0]
|
||||
print(input_L.shape, input_A.shape, input_B.shape)
|
||||
|
||||
ret = []
|
||||
|
||||
if (input_A.shape[0] == batch_size) and (input_B.shape[0]
|
||||
== batch_size):
|
||||
|
||||
for i in range(batch_size):
|
||||
|
||||
input_L_NP = from_torch_image(image=input_L[i])
|
||||
input_A_NP = from_torch_image(image=input_A[i])
|
||||
input_B_NP = from_torch_image(image=input_B[i])
|
||||
|
||||
Y_MAX = input_L_NP.shape[0]
|
||||
X_MAX = input_L_NP.shape[1]
|
||||
|
||||
if (input_A_NP.shape[0]
|
||||
== Y_MAX) and (input_B_NP.shape[0] == Y_MAX) and (
|
||||
(input_A_NP.shape[1] == X_MAX) and
|
||||
(input_B_NP.shape[1] == X_MAX)):
|
||||
|
||||
image = np.zeros((Y_MAX, X_MAX, 3), dtype=np.uint8)
|
||||
|
||||
image[:, :, 0] = input_L_NP
|
||||
image[:, :, 1] = input_A_NP
|
||||
image[:, :, 2] = input_B_NP
|
||||
|
||||
image = cv2.cvtColor(image, cv2.COLOR_LAB2RGB)
|
||||
image = to_torch_image(image).unsqueeze(0)
|
||||
print('image.shape')
|
||||
print(image.shape)
|
||||
ret.append(image)
|
||||
|
||||
else:
|
||||
|
||||
print('Resolution of different layers donot match')
|
||||
|
||||
else:
|
||||
|
||||
print('batch size of different layers donot match')
|
||||
|
||||
ret = torch.cat(ret, dim=0)
|
||||
|
||||
print('ret.shape', ret.shape)
|
||||
|
||||
return (ret, )
|
||||
|
||||
|
||||
class get_mean_and_standard_deviation:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_array": ("MASK", ),
|
||||
"input_mask": ("MASK", ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
"FLOAT",
|
||||
"FLOAT",
|
||||
)
|
||||
|
||||
RETURN_NAMES = (
|
||||
"Mean",
|
||||
"Standard deviation",
|
||||
)
|
||||
|
||||
FUNCTION = "test"
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def test(self, input_array, input_mask):
|
||||
|
||||
input_array = input_array.cpu().numpy()
|
||||
input_mask = input_mask.cpu().numpy()
|
||||
|
||||
mean, sigma = get_mu_sigma(array_input=input_array[0],
|
||||
mask_input=input_mask[0])
|
||||
|
||||
return (
|
||||
mean,
|
||||
sigma,
|
||||
)
|
||||
|
||||
|
||||
class renormalize_array:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_array": ("MASK", ),
|
||||
"input_mask": ("MASK", ),
|
||||
"input_mean": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.001,
|
||||
"round":
|
||||
0.000001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding.
|
||||
"display": "number"
|
||||
}),
|
||||
"input_standard_deviation": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"step": 0.001,
|
||||
"round":
|
||||
0.000001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding.
|
||||
"display": "number"
|
||||
}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK", )
|
||||
RETURN_NAMES = ("Output array as mask", )
|
||||
|
||||
FUNCTION = "test"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "TRI3D"
|
||||
|
||||
def test(
|
||||
self,
|
||||
input_array,
|
||||
input_mask,
|
||||
input_mean,
|
||||
input_standard_deviation,
|
||||
):
|
||||
|
||||
batch_size = input_array.shape[0]
|
||||
|
||||
ret = []
|
||||
|
||||
if input_mask.shape[0] == batch_size:
|
||||
|
||||
for i in range(batch_size):
|
||||
|
||||
tmp = renormalize_array_main(
|
||||
array_input=input_array[i].cpu().numpy(),
|
||||
mask_input=input_mask[i].cpu().numpy(),
|
||||
mu=input_mean,
|
||||
sigma=input_standard_deviation)
|
||||
|
||||
tmp = torch.from_numpy(tmp)
|
||||
tmp = tmp.unsqueeze(0)
|
||||
|
||||
ret.append(tmp)
|
||||
|
||||
else:
|
||||
|
||||
print('batch size of different layers donot match')
|
||||
|
||||
ret = torch.cat(ret, dim=0)
|
||||
|
||||
return (ret, )
|
||||