From c8417fe82cd8cd5298f321be7874aebc4d9f91e4 Mon Sep 17 00:00:00 2001 From: Brandon Thomas Date: Wed, 5 Feb 2025 01:21:27 -0500 Subject: [PATCH] configure model for pose detection --- animation/animate-x/dwpose/wholebody.py | 10 +++++++--- animation/animate-x/process_data.py | 7 ++++--- animation/animate-x/run_process_data.sh | 3 ++- 3 files changed, 13 insertions(+), 7 deletions(-) diff --git a/animation/animate-x/dwpose/wholebody.py b/animation/animate-x/dwpose/wholebody.py index 172f886..9bb29fd 100644 --- a/animation/animate-x/dwpose/wholebody.py +++ b/animation/animate-x/dwpose/wholebody.py @@ -1,3 +1,4 @@ +import os import cv2 import numpy as np @@ -6,12 +7,15 @@ from dwpose.onnxdet import inference_detector from dwpose.onnxpose import inference_pose class Wholebody: - def __init__(self): + def __init__(self, model_directory=None): device = 'cuda' # 'cpu' # providers = ['CPUExecutionProvider' ] if device == 'cpu' else ['CUDAExecutionProvider'] - onnx_det = 'checkpoints/yolox_l.onnx' - onnx_pose = 'checkpoints/dw-ll_ucoco_384.onnx' + + if not model_directory: + model_directory = 'checkpoints' + onnx_det = os.path.join(model_directory, 'yolox_l.onnx') + onnx_pose = os.path.join(model_directory, 'dw-ll_ucoco_384.onnx') self.session_det = ort.InferenceSession(path_or_bytes=onnx_det, providers=providers) self.session_pose = ort.InferenceSession(path_or_bytes=onnx_pose, providers=providers) diff --git a/animation/animate-x/process_data.py b/animation/animate-x/process_data.py index 2c4124b..e307565 100644 --- a/animation/animate-x/process_data.py +++ b/animation/animate-x/process_data.py @@ -89,9 +89,9 @@ def get_logger(name="essmc2"): return logger class DWposeDetector: - def __init__(self): + def __init__(self, model_directory=None): - self.pose_estimation = Wholebody() + self.pose_estimation = Wholebody(model_directory=model_directory) def __call__(self, oriImg): oriImg = oriImg.copy() @@ -235,7 +235,7 @@ def mp_main(args): logger.info("There are {} videos for extracting poses".format(len(video_paths))) logger.info('LOAD: DW Pose Model') - dwpose_model = DWposeDetector() + dwpose_model = DWposeDetector(model_directory=args.model_directory) results_vis = [] for i, file_path in enumerate(video_paths): @@ -336,6 +336,7 @@ if __name__=='__main__': parser.add_argument("--saved_pose_dir", type=str, default="data/saved_pkl",) parser.add_argument("--saved_pose", type=str, default="data/saved_pose",) parser.add_argument("--saved_frame_dir", type=str, default="data/saved_frames",) + parser.add_argument("--model_directory", type=str, default="checkpoints",) args = parser.parse_args() return args diff --git a/animation/animate-x/run_process_data.sh b/animation/animate-x/run_process_data.sh index 4cdba30..01e0856 100755 --- a/animation/animate-x/run_process_data.sh +++ b/animation/animate-x/run_process_data.sh @@ -6,5 +6,6 @@ python process_data.py \ --source_video_paths data/videos/dance_1.mp4 \ --saved_pose_dir data/saved_pkl \ --saved_pose data/saved_pose \ - --saved_frame_dir data/saved_frames + --saved_frame_dir data/saved_frames \ + --model_directory checkpoints2