File size: 2,579 Bytes
cd47a59
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
import cv2
import os
from pathlib import Path
import torch
import numpy as np
from PIL import Image
import sys
import argparse
from tqdm import tqdm
from utils.fs import traverse_folder
import logging
import warnings
os.environ["OPENCV_LOG_LEVEL"] = "FATAL"
warnings.filterwarnings("ignore")

# 在全域強制關閉底層的 stderr 輸出,防止 ONNX C++ 多執行緒偷偷印出 libpng 警告
devnull_fd = open(os.devnull, 'w')
os.dup2(devnull_fd.fileno(), sys.stderr.fileno())

project_root = Path(__file__).parent.parent.parent.parent
sys.path.append(os.path.join(project_root, "DWPose/ControlNet-v1-1-nightly"))
from annotator.dwpose import DWposeDetector
def process_single_image(image_path, detector, output_dir):
    img_name = Path(image_path).name

    out_path = output_dir.joinpath(img_name)
    if os.path.exists(out_path):
        return

    output_dir.mkdir(parents=True, exist_ok=True)

    frame_pil = Image.open(image_path)
    image = cv2.imread(str(image_path))
        
    result = detector(image)
    result = cv2.resize(result, dsize=frame_pil.size, interpolation=cv2.INTER_CUBIC)

    Image.fromarray(result).save(out_path)


def process_batch_images(image_list, detector, output_dir):
    # tqdm 預設印在 stderr,因為 stderr 被我們全域拔掉了,所以強制改印到 stdout
    for i, image_path in enumerate(tqdm(image_list, desc="Generating DWPose", file=sys.stdout)):
        process_single_image(image_path, detector, output_dir)


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--input", type=str, default="", help="image file path or folder include images")
    parser.add_argument("--output", type=str, default="./dwpose", help="Specify output directory")

    args = parser.parse_args()
    image_paths = []

    imgs_path = args.input
    if os.path.isdir(imgs_path):
        for file_path in traverse_folder(imgs_path):
            if os.path.isfile(file_path) and str(file_path).endswith(
                (".jpg", ".png", ".jpeg")
            ):
                image_paths.append(file_path)
    elif imgs_path.suffix in [".jpg", ".png", ".jpeg"]:
        image_paths.append(imgs_path)
    print("Initializing...")
    try:
        detector = DWposeDetector()
        output_dir = Path(args.output)
        if not os.path.exists(output_dir):
            os.makedirs(output_dir)
        process_batch_images(image_paths, detector, output_dir)
    except Exception as e:
        import traceback
        traceback.print_exc(file=sys.stdout)
        sys.exit(1)

    print("finished")