Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
230131d73d | ||
|
|
85e3756721 | ||
|
|
9ccce9efed | ||
|
|
2c1173254e | ||
|
|
abaf432b6f | ||
|
|
a2c78015f4 |
@@ -1,29 +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)
|
||||
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)
|
||||
cloth_seg = generate_mask(img, net=net, device=device)
|
||||
return cloth_seg
|
||||
|
||||
|
||||
INPUT_PATH = "./input/"
|
||||
OUTPUT_PATH = "./output/"
|
||||
|
||||
import os
|
||||
import os
|
||||
for cur_image in 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")
|
||||
|
||||
cv2.imwrite(OUTPUT_PATH + cur_image,
|
||||
cv2.cvtColor(src=cloth_seg, code=cv2.COLOR_RGB2BGR))
|
||||
# cloth_seg.save(OUTPUT_PATH + cur_image, format="PNG")
|
||||
|
||||
@@ -14,6 +14,19 @@ 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):
|
||||
@@ -104,44 +117,56 @@ 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)
|
||||
# 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()
|
||||
|
||||
# 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
|
||||
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 ####')
|
||||
|
||||
# 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")
|
||||
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
|
||||
|
||||
return extracted_garment
|
||||
|
||||
# def generate_mask(input_image, net, device='cpu'):
|
||||
# img = input_image
|
||||
@@ -232,4 +257,4 @@ if __name__ == '__main__':
|
||||
parser.add_argument('--checkpoint_path', type=str, default='model/cloth_segm.pth', help='Path to the checkpoint file')
|
||||
args = parser.parse_args()
|
||||
|
||||
main(args)
|
||||
main(args)
|
||||
|
||||
Reference in New Issue
Block a user