update onnx support for humanparsing

https://github.com/levihsu/OOTDiffusion/commit/c5945e2b603baa24904fbb00dbb6612e809a1142
This commit is contained in:
iyume
2024-03-10 00:41:25 +08:00
parent b7f2132de3
commit 4fbcf14228
6 changed files with 51 additions and 98 deletions
+1 -1
View File
@@ -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):
-14
View File
@@ -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
View File
@@ -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
+22
View File
@@ -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
View File
@@ -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
View File
@@ -1,6 +1,5 @@
torch
torchvision
torchaudio
numpy
scipy
scikit-image
@@ -14,3 +13,4 @@ tqdm
gradio
einops
ninja
onnxruntime