2 changed files with 60 additions and 29 deletions
+12 -6
View File
@@ -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")
+48 -23
View File
@@ -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)