Spaces:
Paused
Paused
File size: 12,551 Bytes
09462dc | 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 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 | import os
import sys
# 动态添加项目根目录到 sys.path,这样就不需要 export PYTHONPATH
current_dir = os.path.dirname(os.path.abspath(__file__))
project_root = os.path.dirname(current_dir) # SCAIL_Pose 目录
if project_root not in sys.path:
sys.path.insert(0, project_root)
import cv2
import torch
import pickle
import torchvision
import shutil
import glob
import random
from tqdm import tqdm
import decord
from decord import VideoReader, cpu, gpu
from torchvision.transforms import ToPILImage
from PIL import Image
import numpy as np
import argparse
from NLFPoseExtract.nlf_render import render_nlf_as_images, collect_smpl_poses, shift_dwpose_according_to_nlf, p3d_single_p2d
from NLFPoseExtract.nlf_draw import intrinsic_matrix_from_field_of_view, process_data_to_COCO_format, p3d_to_p2d
from DWPoseProcess.dwpose import DWposeDetector
from concurrent.futures import ProcessPoolExecutor, as_completed
import multiprocessing
import traceback
from NLFPoseExtract.extract_nlfpose_batch import process_video_nlf
from NLFPoseExtract.reshape_utils_3d import reshapePool3d
try:
import moviepy.editor as mpy
except:
import moviepy as mpy
from torchvision.transforms.functional import center_crop, resize
from torchvision.transforms import InterpolationMode
import torchvision.transforms as TT
import copy
from NLFPoseExtract.align3d import solve_new_camera_params_central, solve_new_camera_params_down
def recollect_nlf(data):
new_data = []
for item in data:
new_item = item.copy()
if len(item['bboxes']) > 0:
new_item['bboxes'] = item['bboxes'][:1]
new_item['nlfpose'] = item['nlfpose'][:1]
new_data.append(new_item)
return new_data
def recollect_dwposes(poses):
new_poses = []
for pose in poses:
new_pose = pose.copy()
for i in range(1):
bodies = pose["bodies"]
faces = pose["faces"][i:i+1]
hands = pose["hands"][2*i:2*i+2]
candidate = bodies["candidate"][i:i+1] # candidate是所有点的坐标和置信度
subset = bodies["subset"][i:i+1] # subset是认为的有效点
new_pose = {
"bodies": {
"candidate": candidate,
"subset": subset
},
"faces": faces,
"hands": hands
}
new_poses.append(new_pose)
return new_poses
def resize_for_rectangle_crop(arr, image_size, reshape_mode='random'):
if arr.shape[3] / arr.shape[2] > image_size[1] / image_size[0]:
arr = resize(arr, size=[image_size[0], int(arr.shape[3] * image_size[0] / arr.shape[2])], interpolation=InterpolationMode.BICUBIC)
else:
arr = resize(arr, size=[int(arr.shape[2] * image_size[1] / arr.shape[3]), image_size[1]], interpolation=InterpolationMode.BICUBIC)
h, w = arr.shape[2], arr.shape[3]
delta_h = h - image_size[0]
delta_w = w - image_size[1]
if reshape_mode == 'random' or reshape_mode == 'none':
top = np.random.randint(0, delta_h + 1)
left = np.random.randint(0, delta_w + 1)
elif reshape_mode == 'center':
top, left = delta_h // 2, delta_w // 2
else:
raise NotImplementedError
arr = TT.functional.crop(
arr, top=top, left=left, height=image_size[0], width=image_size[1]
)
return arr
def scale_faces(poses, pose_2d_ref):
# 输入:两个list of dict,poses[0]['faces'].shape: 1, 68, 2 , poses_ref[0]['faces'].shape: 1, 68, 2
# 根据脸部的中心点,对poses中的脸部关键点进行缩放
# 也即:计算ref里面脸部中心点(idx: 30)到其他脸部关键点的中心距离, 计算poses里面脸部中心点到其他脸部关键点的中心距离,得到scale_n
# 对scale_n 取一下0.8-1.5的上下界,然后应用在poses上
# 注意:需要inplace改变poses
ref = pose_2d_ref[0]
pose_0 = poses[0]
face_0 = pose_0['faces'] # shape: (1, 68, 2)
face_ref = ref['faces']
# 提取 numpy 数组
face_0 = np.array(face_0[0]) # (68, 2)
face_ref = np.array(face_ref[0])
# 中心点(鼻尖或面部中心)
center_idx = 30
center_0 = face_0[center_idx]
center_ref = face_ref[center_idx]
# 计算到中心点的距离
dist = np.linalg.norm(face_0 - center_0, axis=1)
dist_ref = np.linalg.norm(face_ref - center_ref, axis=1)
# 避免中心点自身的 0 距离影响
dist = np.delete(dist, center_idx)
dist_ref = np.delete(dist_ref, center_idx)
mean_dist = np.mean(dist)
mean_dist_ref = np.mean(dist_ref)
if mean_dist < 1e-6:
scale_n = 1.0
else:
scale_n = mean_dist_ref / mean_dist
# 限制在 [0.8, 1.5]
scale_n = np.clip(scale_n, 0.8, 1.5)
for i, pose in enumerate(poses):
face = pose['faces']
# 提取 numpy 数组
face = np.array(face[0]) # (68, 2)
center = face[center_idx]
scaled_face = (face - center) * scale_n + center
poses[i]['faces'][0] = scaled_face
body = pose['bodies']
candidate = body['candidate']
candidate_np = np.array(candidate[0]) # (14, 2)
body_center = candidate_np[0]
scaled_candidate = (candidate_np - body_center) * scale_n + body_center
poses[i]['bodies']['candidate'][0] = scaled_candidate
# inplace 修改
pose['faces'][0] = scaled_face
return scale_n
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Process video with NLF pose estimation')
parser.add_argument('--subdir', type=str, default="../examples/001", help='Path to the subdirectory to process')
parser.add_argument('--model_path', type=str, default='pretrained_weights/nlf_l_multi_0.3.2.torchscript',
help='Path to NLF model')
parser.add_argument('--use_align', action='store_true', help='Whether to use 2D keypoints from reference image for alignment')
parser.add_argument('--resolution', type=int, nargs=2, default=[512, 896],
metavar=('HEIGHT', 'WIDTH'),
help='Target resolution as [height, width], currently only [512, 896] are supported')
args = parser.parse_args()
subdir = args.subdir
model_nlf = torch.jit.load(args.model_path).cuda().eval()
decord.bridge.set_bridge("torch")
# 设置路径
mp4_path = os.path.join(subdir, 'driving.mp4')
if not os.path.exists(mp4_path):
raise FileNotFoundError(f"No video file found in {subdir}")
if args.use_align:
out_path_aligned = os.path.join(subdir, 'rendered_aligned.mp4')
else:
out_path_aligned = os.path.join(subdir, 'rendered.mp4')
ref_image_path = os.path.join(subdir, 'ref_image.jpg')
if not os.path.exists(ref_image_path):
ref_image_path = os.path.join(subdir, 'ref_image.png')
if not os.path.exists(ref_image_path):
ref_image_path = os.path.join(subdir, 'ref.jpg')
if not os.path.exists(ref_image_path):
raise FileNotFoundError(f"No reference image found in {subdir}")
print(f"Processing: {subdir}")
print(f"Video: {mp4_path}")
print(f"Reference: {ref_image_path}")
print(f"Resolution: {args.resolution}")
# 读取视频
vr = VideoReader(mp4_path)
vr_frames = vr.get_batch(list(range(len(vr)))) # T H W C
sampling_image_size = args.resolution
if vr_frames.shape[1] < vr_frames.shape[2]:
target_H, target_W = sampling_image_size
else:
target_W, target_H = sampling_image_size
vr_frames = resize_for_rectangle_crop(vr_frames.permute(0, 3, 1, 2), [target_H, target_W], reshape_mode='center').permute(0, 2, 3, 1) # T H W C ->T C H W -> T H W C
# 读取参考图片
img_ref = cv2.imread(ref_image_path)
img_ref = cv2.cvtColor(img_ref, cv2.COLOR_BGR2RGB)
vr_frames_ref = torch.from_numpy(img_ref).unsqueeze(0)
vr_frames_ref = resize_for_rectangle_crop(vr_frames_ref.permute(0, 3, 1, 2), [target_H, target_W], reshape_mode='center').permute(0, 2, 3, 1) # 1 H W C ->1 C H W -> 1 H W C
# 初始化检测器
detector = DWposeDetector(use_batch=False).to(0)
# 处理Driving视频
print("Processing driving video...")
detector_return_list = []
pil_frames = []
for i in tqdm(range(len(vr_frames)), desc="Detecting poses in video"):
pil_frame = Image.fromarray(vr_frames[i].numpy())
pil_frames.append(pil_frame)
detector_result = detector(pil_frame)
detector_return_list.append(detector_result)
W, H = pil_frames[0].size
poses, scores, det_results = zip(*detector_return_list)
print("Running NLF on driving video...")
nlf_results = process_video_nlf(model_nlf, vr_frames, det_results)
# 处理ref图片
print("Processing reference image...")
detector_return_list_ref = []
pil_frames_ref = []
for i in range(len(vr_frames_ref)):
pil_frame = Image.fromarray(vr_frames_ref[i].numpy())
pil_frames_ref.append(pil_frame)
detector_result = detector(pil_frame)
detector_return_list_ref.append(detector_result)
poses_ref, scores_ref, det_results_ref = zip(*detector_return_list_ref)
print("Running NLF on reference image...")
nlf_results_ref = process_video_nlf(model_nlf, vr_frames_ref, det_results_ref)
# 进行对齐和渲染
print("Aligning and rendering...")
ori_camera_pose = intrinsic_matrix_from_field_of_view([target_H, target_W])
ori_focal = ori_camera_pose[0, 0]
pose_3d_first_driving_frame = nlf_results[0]['nlfpose'][0][0].cpu().numpy() # 3D点 frame-idx bbox-idx detect-idx
pose_3d_coco_first_driving_frame = process_data_to_COCO_format(pose_3d_first_driving_frame)
if args.use_align:
poses_2d_ref = poses_ref[0]['bodies']['candidate'][0][:14]
poses_2d_ref[:, 0] = poses_2d_ref[:, 0] * target_W
poses_2d_ref[:, 1] = poses_2d_ref[:, 1] * target_H
poses_2d_subset = poses_ref[0]['bodies']['subset'][0][:14]
pose_3d_coco_first_driving_frame = pose_3d_coco_first_driving_frame[:14]
valid_indices = []
valid_upper_indices = []
valid_lower_indices = []
upper_body_indices = [0, 2, 3, 5, 6]
lower_body_indices = [9, 10, 12, 13]
excluded_indices = [3, 4, 6, 7] # 去除手
for i in range(len(poses_2d_subset)):
if poses_2d_subset[i] != -1.0 and np.sum(pose_3d_coco_first_driving_frame[i]) != 0:
if i in upper_body_indices:
valid_upper_indices.append(i)
if i in lower_body_indices:
valid_lower_indices.append(i)
if len(valid_lower_indices) >= 4:
print("Align feet")
valid_indices = [1] + valid_lower_indices
else:
print("Align body")
valid_indices = [1] + valid_upper_indices
pose_2d_ref = poses_2d_ref[valid_indices]
pose_3d_coco_first_driving_frame = pose_3d_coco_first_driving_frame[valid_indices]
if len(valid_lower_indices) >= 4:
new_camera_intrinsics, scale_m, scale_s = solve_new_camera_params_down(pose_3d_coco_first_driving_frame, ori_focal, [target_H, target_W], pose_2d_ref)
else:
new_camera_intrinsics, scale_m, scale_s = solve_new_camera_params_central(pose_3d_coco_first_driving_frame, ori_focal, [target_H, target_W], pose_2d_ref)
# m 代表缩放了多少
scale_face = scale_faces(list(poses), list(poses_ref)) # poses[0]['faces'].shape: 1, 68, 2 , poses_ref[0]['faces'].shape: 1, 68, 2
print(f"Scale - m: {scale_m}, face: {scale_face}")
nlf_results = recollect_nlf(nlf_results)
poses = recollect_dwposes(list(poses))
shift_dwpose_according_to_nlf(collect_smpl_poses(nlf_results), poses, ori_camera_pose, new_camera_intrinsics, target_H, target_W, scale_x=scale_m, scale_y=scale_m*scale_s)
print("Rendering final video...")
frames_np = render_nlf_as_images(nlf_results, poses, reshape_pool=None, intrinsic_matrix=new_camera_intrinsics)
else:
nlf_results = recollect_nlf(nlf_results)
print("Rendering final video...")
frames_np = render_nlf_as_images(nlf_results, poses, reshape_pool=None, intrinsic_matrix=ori_camera_pose)
mpy.ImageSequenceClip(frames_np, fps=16).write_videofile(out_path_aligned)
print(f"Done! Output saved to: {out_path_aligned}")
|