mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
configure model for pose detection
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user