update onnx support for humanparsing
https://github.com/levihsu/OOTDiffusion/commit/c5945e2b603baa24904fbb00dbb6612e809a1142
This commit is contained in:
+1
-1
@@ -63,7 +63,7 @@ class LoadOOTDPipelineHub(LoadOOTDPipeline):
|
||||
display_name = "Load OOTDiffusion from Hub🤗"
|
||||
|
||||
repo_id = "levihsu/OOTDiffusion"
|
||||
repo_revision = "6150dd90d7302f21348b8d8461766a400c00716a"
|
||||
repo_revision = "d33c517dc1b0718ea1136533e3720bb08fae641b"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
|
||||
@@ -1,14 +0,0 @@
|
||||
from .parsing_api import load_atr_model, load_lip_model, inference
|
||||
|
||||
|
||||
class Parsing:
|
||||
def __init__(self, atr_model_path, lip_model_path, *, device):
|
||||
self.device = device
|
||||
self.atr_model = load_atr_model(atr_model_path).to(device)
|
||||
self.lip_model = load_lip_model(lip_model_path).to(device)
|
||||
|
||||
def __call__(self, input_image):
|
||||
parsed_image, face_mask = inference(
|
||||
self.atr_model, self.lip_model, input_image, device=self.device
|
||||
)
|
||||
return parsed_image, face_mask
|
||||
+24
-78
@@ -1,11 +1,6 @@
|
||||
from pathlib import Path
|
||||
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
import cv2
|
||||
from . import networks
|
||||
from collections import OrderedDict
|
||||
import torchvision.transforms as transforms
|
||||
from torch.utils.data import DataLoader
|
||||
from .datasets.simple_extractor_dataset import SimpleFolderDataset
|
||||
@@ -116,45 +111,7 @@ def refine_hole(parsing_result_filled, parsing_result, arm_mask):
|
||||
cv2.drawContours(refine_hole_mask, contours, i, color=255, thickness=-1)
|
||||
return refine_hole_mask + arm_mask
|
||||
|
||||
|
||||
|
||||
def load_atr_model(path: str):
|
||||
# load atr model
|
||||
num_classes = 18
|
||||
label = ['Background', 'Hat', 'Hair', 'Sunglasses', 'Upper-clothes', 'Skirt', 'Pants', 'Dress', 'Belt',
|
||||
'Left-shoe', 'Right-shoe', 'Face', 'Left-leg', 'Right-leg', 'Left-arm', 'Right-arm', 'Bag', 'Scarf']
|
||||
model = networks.init_model('resnet101', num_classes=num_classes, pretrained=None)
|
||||
state_dict = torch.load(path)['state_dict']
|
||||
new_state_dict = OrderedDict()
|
||||
for k, v in state_dict.items():
|
||||
name = k[7:] # remove `module.`
|
||||
new_state_dict[name] = v
|
||||
model.load_state_dict(new_state_dict)
|
||||
model.cuda()
|
||||
model.eval()
|
||||
# load lip model
|
||||
return model
|
||||
|
||||
def load_lip_model(path: str):
|
||||
# load atr model
|
||||
num_classes = 20
|
||||
label = ['Background', 'Hat', 'Hair', 'Glove', 'Sunglasses', 'Upper-clothes', 'Dress', 'Coat',
|
||||
'Socks', 'Pants', 'Jumpsuits', 'Scarf', 'Skirt', 'Face', 'Left-arm', 'Right-arm',
|
||||
'Left-leg', 'Right-leg', 'Left-shoe', 'Right-shoe']
|
||||
model = networks.init_model('resnet101', num_classes=num_classes, pretrained=None)
|
||||
state_dict = torch.load(path)['state_dict']
|
||||
new_state_dict = OrderedDict()
|
||||
for k, v in state_dict.items():
|
||||
name = k[7:] # remove `module.`
|
||||
new_state_dict[name] = v
|
||||
model.load_state_dict(new_state_dict)
|
||||
model.cuda()
|
||||
model.eval()
|
||||
# load lip model
|
||||
return model
|
||||
|
||||
def inference(model, lip_model, input_dir, *, device):
|
||||
# load datasetloader
|
||||
def onnx_inference(session, lip_session, input_dir):
|
||||
transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=[0.406, 0.456, 0.485], std=[0.225, 0.224, 0.229])
|
||||
@@ -168,19 +125,14 @@ def inference(model, lip_model, input_dir, *, device):
|
||||
s = meta['scale'].numpy()[0]
|
||||
w = meta['width'].numpy()[0]
|
||||
h = meta['height'].numpy()[0]
|
||||
|
||||
output = model(image.to(device))
|
||||
|
||||
output = session.run(None, {"input.1": image.numpy().astype(np.float32)})
|
||||
upsample = torch.nn.Upsample(size=[512, 512], mode='bilinear', align_corners=True)
|
||||
upsample_output = upsample(output[0][-1][0].unsqueeze(0))
|
||||
upsample_output = upsample(torch.from_numpy(output[1][0]).unsqueeze(0))
|
||||
upsample_output = upsample_output.squeeze()
|
||||
upsample_output = upsample_output.permute(1, 2, 0) # CHW -> HWC
|
||||
logits_result = transform_logits(upsample_output.data.cpu().numpy(), c, s, w, h, input_size=[512, 512])
|
||||
|
||||
# delete irregular classes, e.g. pants/ skirts over clothes
|
||||
parsing_result = np.argmax(logits_result, axis=2)
|
||||
parsing_result = np.pad(parsing_result, pad_width=1, mode='constant', constant_values=0)
|
||||
|
||||
# try holefilling the clothes part
|
||||
arm_mask = (parsing_result == 14).astype(np.float32) \
|
||||
+ (parsing_result == 15).astype(np.float32)
|
||||
@@ -189,7 +141,6 @@ def inference(model, lip_model, input_dir, *, device):
|
||||
dst = hole_fill(img.astype(np.uint8))
|
||||
parsing_result_filled = dst / 255 * 4
|
||||
parsing_result_woarm = np.where(parsing_result_filled == 4, parsing_result_filled, parsing_result)
|
||||
|
||||
# add back arm and refined hole between arm and cloth
|
||||
refine_hole_mask = refine_hole(parsing_result_filled.astype(np.uint8), parsing_result.astype(np.uint8),
|
||||
arm_mask.astype(np.uint8))
|
||||
@@ -197,38 +148,33 @@ def inference(model, lip_model, input_dir, *, device):
|
||||
# remove padding
|
||||
parsing_result = parsing_result[1:-1, 1:-1]
|
||||
|
||||
dataset_lip = SimpleFolderDataset(root=input_dir, input_size=[473, 473], transform=transform)
|
||||
dataloader_lip = DataLoader(dataset_lip)
|
||||
with torch.no_grad():
|
||||
for _, batch in enumerate(tqdm(dataloader_lip)):
|
||||
image, meta = batch
|
||||
c = meta['center'].numpy()[0]
|
||||
s = meta['scale'].numpy()[0]
|
||||
w = meta['width'].numpy()[0]
|
||||
h = meta['height'].numpy()[0]
|
||||
|
||||
dataset_lip = SimpleFolderDataset(root=input_dir, input_size=[473, 473], transform=transform)
|
||||
dataloader_lip = DataLoader(dataset_lip)
|
||||
with torch.no_grad():
|
||||
for _, batch in enumerate(tqdm(dataloader_lip)):
|
||||
image, meta = batch
|
||||
c = meta['center'].numpy()[0]
|
||||
s = meta['scale'].numpy()[0]
|
||||
w = meta['width'].numpy()[0]
|
||||
h = meta['height'].numpy()[0]
|
||||
|
||||
output_lip = lip_model(image.to(device))
|
||||
|
||||
upsample = torch.nn.Upsample(size=[473, 473], mode='bilinear', align_corners=True)
|
||||
upsample_output_lip = upsample(output_lip[0][-1][0].unsqueeze(0))
|
||||
upsample_output_lip = upsample_output_lip.squeeze()
|
||||
upsample_output_lip = upsample_output_lip.permute(1, 2, 0) # CHW -> HWC
|
||||
logits_result_lip = transform_logits(upsample_output_lip.data.cpu().numpy(), c, s, w, h, input_size=[473, 473])
|
||||
parsing_result_lip = np.argmax(logits_result_lip, axis=2)
|
||||
output_lip = lip_session.run(None, {"input.1": image.numpy().astype(np.float32)})
|
||||
upsample = torch.nn.Upsample(size=[473, 473], mode='bilinear', align_corners=True)
|
||||
upsample_output_lip = upsample(torch.from_numpy(output_lip[1][0]).unsqueeze(0))
|
||||
upsample_output_lip = upsample_output_lip.squeeze()
|
||||
upsample_output_lip = upsample_output_lip.permute(1, 2, 0) # CHW -> HWC
|
||||
logits_result_lip = transform_logits(upsample_output_lip.data.cpu().numpy(), c, s, w, h,
|
||||
input_size=[473, 473])
|
||||
parsing_result_lip = np.argmax(logits_result_lip, axis=2)
|
||||
# add neck parsing result
|
||||
neck_mask = np.logical_and(np.logical_not((parsing_result_lip == 13).astype(np.float32)), (parsing_result == 11).astype(np.float32))
|
||||
# filter out small part of neck
|
||||
neck_mask = refine_mask(neck_mask)
|
||||
# Image.fromarray(((neck_mask > 0) * 127.5 + 127.5).astype(np.uint8)).save("neck_mask.jpg")
|
||||
neck_mask = np.logical_and(np.logical_not((parsing_result_lip == 13).astype(np.float32)),
|
||||
(parsing_result == 11).astype(np.float32))
|
||||
parsing_result = np.where(neck_mask, 18, parsing_result)
|
||||
palette = get_palette(19)
|
||||
parsing_result_path = os.path.join('parsed.png')
|
||||
output_img = Image.fromarray(np.asarray(parsing_result, dtype=np.uint8))
|
||||
output_img.putpalette(palette)
|
||||
# output_img.save(parsing_result_path)
|
||||
face_mask = torch.from_numpy((parsing_result == 11).astype(np.float32))
|
||||
|
||||
face_mask = torch.from_numpy((parsing_result == 11).astype(np.float32))
|
||||
|
||||
return output_img, face_mask
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
from pathlib import Path
|
||||
import os
|
||||
import onnxruntime as ort
|
||||
from .parsing_api import onnx_inference
|
||||
import torch
|
||||
|
||||
|
||||
class Parsing:
|
||||
def __init__(self, atr_model_path, lip_model_path):
|
||||
session_options = ort.SessionOptions()
|
||||
session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
||||
session_options.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
|
||||
# session_options.add_session_config_entry('gpu_id', str(gpu_id))
|
||||
self.session = ort.InferenceSession(atr_model_path,
|
||||
sess_options=session_options, providers=['CPUExecutionProvider'])
|
||||
self.lip_session = ort.InferenceSession(lip_model_path,
|
||||
sess_options=session_options, providers=['CPUExecutionProvider'])
|
||||
|
||||
|
||||
def __call__(self, input_image):
|
||||
parsed_image, face_mask = onnx_inference(self.session, self.lip_session, input_image)
|
||||
return parsed_image, face_mask
|
||||
+3
-4
@@ -14,7 +14,7 @@ from transformers import (
|
||||
)
|
||||
|
||||
from . import pipelines_ootd
|
||||
from .humanparsing.aigc_run_parsing import Parsing
|
||||
from .humanparsing.run_parsing import Parsing
|
||||
from .openpose.run_openpose import OpenPose
|
||||
|
||||
#! Necessary for OotdPipeline.from_pretrained
|
||||
@@ -46,15 +46,14 @@ class OOTDiffusion:
|
||||
UNET_PATH = MODEL_PATH / "ootd_dc" / "checkpoint-36000"
|
||||
|
||||
atr_model_path = (
|
||||
Path(root) / "checkpoints/humanparsing/exp-schp-201908301523-atr.pth"
|
||||
Path(root) / "checkpoints/humanparsing/parsing_atr.onnx"
|
||||
)
|
||||
lip_model_path = (
|
||||
Path(root) / "checkpoints/humanparsing/exp-schp-201908261155-lip.pth"
|
||||
Path(root) / "checkpoints/humanparsing/parsing_lip.onnx"
|
||||
)
|
||||
self.parsing_model = Parsing(
|
||||
atr_model_path=str(atr_model_path),
|
||||
lip_model_path=str(lip_model_path),
|
||||
device=self.device,
|
||||
)
|
||||
body_pose_model_path = (
|
||||
Path(root) / "checkpoints/openpose/ckpts/body_pose_model.pth"
|
||||
|
||||
+1
-1
@@ -1,6 +1,5 @@
|
||||
torch
|
||||
torchvision
|
||||
torchaudio
|
||||
numpy
|
||||
scipy
|
||||
scikit-image
|
||||
@@ -14,3 +13,4 @@ tqdm
|
||||
gradio
|
||||
einops
|
||||
ninja
|
||||
onnxruntime
|
||||
|
||||
Reference in New Issue
Block a user