cnywt/SyncTalk
0
1import os2import glob3import tqdm4import json5import argparse6import cv27import numpy as np8import torch9import torch.nn.functional as F10import face_alignment11from face_tracking.util import euler2rot12 13 14def extract_audio(path, out_path, sample_rate=16000):15 16 print(f'[INFO] ===== extract audio from {path} to {out_path} =====')17 cmd = f'ffmpeg -i {path} -f wav -ar {sample_rate} {out_path}'18 os.system(cmd)19 print(f'[INFO] ===== extracted audio =====')20 21def extract_audio_features(path, mode='ave'):22 23 print(f'[INFO] ===== extract audio labels for {path} =====')24 if mode == 'ave':25 print(f'AVE has been integrated into the training code, no need to extract audio features')26 elif mode == "deepspeech": # deepspeech27 cmd = f'python data_utils/deepspeech_features/extract_ds_features.py --input {path}'28 os.system(cmd)29 elif mode == 'hubert':30 cmd = f'python data_utils/hubert.py --wav {path}' # save to data/<name>_hu.npy31 os.system(cmd)32 print(f'[INFO] ===== extracted audio labels =====')33 34 35def extract_images(path, out_path, fps=25):36 37 print(f'[INFO] ===== extract images from {path} to {out_path} =====')38 cmd = f'ffmpeg -i {path} -vf fps={fps} -qmin 1 -q:v 1 -start_number 0 {os.path.join(out_path, "%d.jpg")}'39 os.system(cmd)40 print(f'[INFO] ===== extracted images =====')41 42 43def extract_semantics(ori_imgs_dir, parsing_dir):44 45 print(f'[INFO] ===== extract semantics from {ori_imgs_dir} to {parsing_dir} =====')46 cmd = f'python data_utils/face_parsing/test.py --respath={parsing_dir} --imgpath={ori_imgs_dir}'47 os.system(cmd)48 print(f'[INFO] ===== extracted semantics =====')49 50 51def extract_landmarks(ori_imgs_dir):52 53 print(f'[INFO] ===== extract face landmarks from {ori_imgs_dir} =====')54 try:55 fa = face_alignment.FaceAlignment(face_alignment.LandmarksType._2D, flip_input=False)56 except:57 fa = face_alignment.FaceAlignment(face_alignment.LandmarksType.TWO_D, flip_input=False)58 image_paths = glob.glob(os.path.join(ori_imgs_dir, '*.jpg'))59 for image_path in tqdm.tqdm(image_paths):60 input = cv2.imread(image_path, cv2.IMREAD_UNCHANGED) # [H, W, 3]61 input = cv2.cvtColor(input, cv2.COLOR_BGR2RGB)62 preds = fa.get_landmarks(input)63 if len(preds) > 0:64 lands = preds[0].reshape(-1, 2)[:,:2]65 np.savetxt(image_path.replace('jpg', 'lms'), lands, '%f')66 del fa67 print(f'[INFO] ===== extracted face landmarks =====')68 69 70def extract_background(base_dir, ori_imgs_dir):71 72 print(f'[INFO] ===== extract background image from {ori_imgs_dir} =====')73 74 from sklearn.neighbors import NearestNeighbors75 76 image_paths = glob.glob(os.path.join(ori_imgs_dir, '*.jpg'))77 # only use 1/20 image_paths 78 image_paths = image_paths[::20]79 # read one image to get H/W80 tmp_image = cv2.imread(image_paths[0], cv2.IMREAD_UNCHANGED) # [H, W, 3]81 h, w = tmp_image.shape[:2]82 83 # nearest neighbors84 all_xys = np.mgrid[0:h, 0:w].reshape(2, -1).transpose()85 distss = []86 for image_path in tqdm.tqdm(image_paths):87 parse_img = cv2.imread(image_path.replace('ori_imgs', 'parsing').replace('.jpg', '.png'))88 bg = (parse_img[..., 0] == 255) & (parse_img[..., 1] == 255) & (parse_img[..., 2] == 255)89 fg_xys = np.stack(np.nonzero(~bg)).transpose(1, 0)90 nbrs = NearestNeighbors(n_neighbors=1, algorithm='kd_tree').fit(fg_xys)91 dists, _ = nbrs.kneighbors(all_xys)92 distss.append(dists)93 94 distss = np.stack(distss)95 max_dist = np.max(distss, 0)96 max_id = np.argmax(distss, 0)97 98 bc_pixs = max_dist > 599 bc_pixs_id = np.nonzero(bc_pixs)100 bc_ids = max_id[bc_pixs]101 102 imgs = []103 num_pixs = distss.shape[1]104 for image_path in image_paths:105 img = cv2.imread(image_path)106 imgs.append(img)107 imgs = np.stack(imgs).reshape(-1, num_pixs, 3)108 109 bc_img = np.zeros((h*w, 3), dtype=np.uint8)110 bc_img[bc_pixs_id, :] = imgs[bc_ids, bc_pixs_id, :]111 bc_img = bc_img.reshape(h, w, 3)112 113 max_dist = max_dist.reshape(h, w)114 bc_pixs = max_dist > 5115 bg_xys = np.stack(np.nonzero(~bc_pixs)).transpose()116 fg_xys = np.stack(np.nonzero(bc_pixs)).transpose()117 nbrs = NearestNeighbors(n_neighbors=1, algorithm='kd_tree').fit(fg_xys)118 distances, indices = nbrs.kneighbors(bg_xys)119 bg_fg_xys = fg_xys[indices[:, 0]]120 bc_img[bg_xys[:, 0], bg_xys[:, 1], :] = bc_img[bg_fg_xys[:, 0], bg_fg_xys[:, 1], :]121 122 cv2.imwrite(os.path.join(base_dir, 'bc.jpg'), bc_img)123 124 print(f'[INFO] ===== extracted background image =====')125 126 127def extract_torso_and_gt(base_dir, ori_imgs_dir):128 129 print(f'[INFO] ===== extract torso and gt images for {base_dir} =====')130 131 from scipy.ndimage import binary_erosion, binary_dilation132 133 # load bg134 bg_image = cv2.imread(os.path.join(base_dir, 'bc.jpg'), cv2.IMREAD_UNCHANGED)135 136 image_paths = glob.glob(os.path.join(ori_imgs_dir, '*.jpg'))137 138 for image_path in tqdm.tqdm(image_paths):139 # read ori image140 ori_image = cv2.imread(image_path, cv2.IMREAD_UNCHANGED) # [H, W, 3]141 142 # read semantics143 seg = cv2.imread(image_path.replace('ori_imgs', 'parsing').replace('.jpg', '.png'))144 mask_img = np.zeros_like(seg)145 head_part = (seg[..., 0] == 255) & (seg[..., 1] == 0) & (seg[..., 2] == 0)146 neck_part = (seg[..., 0] == 0) & (seg[..., 1] == 255) & (seg[..., 2] == 0)147 torso_part = (seg[..., 0] == 0) & (seg[..., 1] == 0) & (seg[..., 2] == 255)148 bg_part = (seg[..., 0] == 255) & (seg[..., 1] == 255) & (seg[..., 2] == 255)149 mask_img[head_part, :] = 255150 cv2.imwrite(image_path.replace('ori_imgs', 'face_mask').replace('.jpg', '.png'), mask_img)151 # get gt image152 gt_image = ori_image.copy()153 gt_image[bg_part] = bg_image[bg_part]154 cv2.imwrite(image_path.replace('ori_imgs', 'gt_imgs'), gt_image)155 156 # get torso image157 torso_image = gt_image.copy() # rgb158 torso_image[head_part] = bg_image[head_part]159 torso_alpha = 255 * np.ones((gt_image.shape[0], gt_image.shape[1], 1), dtype=np.uint8) # alpha160 161 # torso part "vertical" in-painting...162 L = 8 + 1163 torso_coords = np.stack(np.nonzero(torso_part), axis=-1) # [M, 2]164 # lexsort: sort 2D coords first by y then by x, 165 # ref: https://stackoverflow.com/questions/2706605/sorting-a-2d-numpy-array-by-multiple-axes166 inds = np.lexsort((torso_coords[:, 0], torso_coords[:, 1]))167 torso_coords = torso_coords[inds]168 # choose the top pixel for each column169 u, uid, ucnt = np.unique(torso_coords[:, 1], return_index=True, return_counts=True)170 top_torso_coords = torso_coords[uid] # [m, 2]171 # only keep top-is-head pixels172 top_torso_coords_up = top_torso_coords.copy() - np.array([1, 0])173 mask = head_part[tuple(top_torso_coords_up.T)] 174 if mask.any():175 top_torso_coords = top_torso_coords[mask]176 # get the color177 top_torso_colors = gt_image[tuple(top_torso_coords.T)] # [m, 3]178 # construct inpaint coords (vertically up, or minus in x)179 inpaint_torso_coords = top_torso_coords[None].repeat(L, 0) # [L, m, 2]180 inpaint_offsets = np.stack([-np.arange(L), np.zeros(L, dtype=np.int32)], axis=-1)[:, None] # [L, 1, 2]181 inpaint_torso_coords += inpaint_offsets182 inpaint_torso_coords = inpaint_torso_coords.reshape(-1, 2) # [Lm, 2]183 inpaint_torso_colors = top_torso_colors[None].repeat(L, 0) # [L, m, 3]184 darken_scaler = 0.98 ** np.arange(L).reshape(L, 1, 1) # [L, 1, 1]185 inpaint_torso_colors = (inpaint_torso_colors * darken_scaler).reshape(-1, 3) # [Lm, 3]186 # set color187 torso_image[tuple(inpaint_torso_coords.T)] = inpaint_torso_colors188 189 inpaint_torso_mask = np.zeros_like(torso_image[..., 0]).astype(bool)190 inpaint_torso_mask[tuple(inpaint_torso_coords.T)] = True191 else:192 inpaint_torso_mask = None193 194 push_down = 4195 L = 48 + push_down + 1196 197 neck_part = binary_dilation(neck_part, structure=np.array([[0, 1, 0], [0, 1, 0], [0, 1, 0]], dtype=bool), iterations=3)198 199 neck_coords = np.stack(np.nonzero(neck_part), axis=-1) # [M, 2]200 inds = np.lexsort((neck_coords[:, 0], neck_coords[:, 1]))201 neck_coords = neck_coords[inds]202 u, uid, ucnt = np.unique(neck_coords[:, 1], return_index=True, return_counts=True)203 top_neck_coords = neck_coords[uid] # [m, 2]204 top_neck_coords_up = top_neck_coords.copy() - np.array([1, 0])205 mask = head_part[tuple(top_neck_coords_up.T)] 206 207 top_neck_coords = top_neck_coords[mask]208 offset_down = np.minimum(ucnt[mask] - 1, push_down)209 top_neck_coords += np.stack([offset_down, np.zeros_like(offset_down)], axis=-1)210 # get the color211 top_neck_colors = gt_image[tuple(top_neck_coords.T)] # [m, 3]212 213 # construct inpaint coords (vertically up, or minus in x)214 inpaint_neck_coords = top_neck_coords[None].repeat(L, 0) # [L, m, 2]215 inpaint_offsets = np.stack([-np.arange(L), np.zeros(L, dtype=np.int32)], axis=-1)[:, None] # [L, 1, 2]216 inpaint_neck_coords += inpaint_offsets217 inpaint_neck_coords = inpaint_neck_coords.reshape(-1, 2) # [Lm, 2]218 219 #add220 neck_avg_color = np.mean(gt_image[neck_part], axis=0)221 inpaint_neck_colors = top_neck_colors[None].repeat(L, 0) # [L, m, 3]222 alpha_values = np.linspace(1, 0, L).reshape(L, 1, 1) # [L, 1, 1]223 inpaint_neck_colors = inpaint_neck_colors * alpha_values + neck_avg_color * (1 - alpha_values)224 inpaint_neck_colors = inpaint_neck_colors.reshape(-1, 3) # [Lm, 3]225 torso_image[tuple(inpaint_neck_coords.T)] = inpaint_neck_colors226 227 inpaint_mask = np.zeros_like(torso_image[..., 0]).astype(bool)228 inpaint_mask[tuple(inpaint_neck_coords.T)] = True229 230 blur_img = torso_image.copy()231 blur_img = cv2.GaussianBlur(blur_img, (5, 5), cv2.BORDER_DEFAULT)232 233 torso_image[inpaint_mask] = blur_img[inpaint_mask]234 235 # set mask236 mask = (neck_part | torso_part | inpaint_mask)237 if inpaint_torso_mask is not None:238 mask = mask | inpaint_torso_mask239 torso_image[~mask] = 0240 torso_alpha[~mask] = 0241 242 cv2.imwrite(image_path.replace('ori_imgs', 'torso_imgs').replace('.jpg', '.png'), np.concatenate([torso_image, torso_alpha], axis=-1))243 print(f'[INFO] ===== extracted torso and gt images =====')244 245 246def face_tracking(ori_imgs_dir):247 248 print(f'[INFO] ===== perform face tracking =====')249 250 image_paths = glob.glob(os.path.join(ori_imgs_dir, '*.jpg'))251 252 # read one image to get H/W253 tmp_image = cv2.imread(image_paths[0], cv2.IMREAD_UNCHANGED) # [H, W, 3]254 h, w = tmp_image.shape[:2]255 256 cmd = f'python data_utils/face_tracking/face_tracker.py --path={ori_imgs_dir} --img_h={h} --img_w={w} --frame_num={len(image_paths)}'257 258 os.system(cmd)259 260 print(f'[INFO] ===== finished face tracking =====')261 262# ref: https://github.com/ShunyuYao/DFA-NeRF263def extract_flow(base_dir,ori_imgs_dir,mask_dir, flow_dir):264 print(f'[INFO] ===== extract flow =====')265 torch.cuda.empty_cache()266 ref_id = 2267 image_paths = glob.glob(os.path.join(ori_imgs_dir, '*.jpg'))268 tmp_image = cv2.imread(image_paths[0], cv2.IMREAD_UNCHANGED) # [H, W, 3]269 h, w = tmp_image.shape[:2]270 valid_img_ids = []271 for i in range(100000):272 if os.path.isfile(os.path.join(ori_imgs_dir, '{:d}.lms'.format(i))):273 valid_img_ids.append(i)274 valid_img_num = len(valid_img_ids)275 with open(os.path.join(base_dir, 'flow_list.txt'), 'w') as file:276 for i in range(0, valid_img_num):277 file.write(base_dir + '/ori_imgs/' + '{:d}.jpg '.format(ref_id) +278 base_dir + '/face_mask/' + '{:d}.png '.format(ref_id) +279 base_dir + '/ori_imgs/' + '{:d}.jpg '.format(i) +280 base_dir + '/face_mask/' + '{:d}.png\n'.format(i))281 file.close()282 ext_flow_cmd = 'python data_utils/UNFaceFlow/test_flow.py --datapath=' + base_dir + '/flow_list.txt ' + \283 '--savepath=' + base_dir + '/flow_result' + \284 ' --width=' + str(w) + ' --height=' + str(h)285 os.system(ext_flow_cmd)286 face_img = cv2.imread(os.path.join(ori_imgs_dir, '{:d}.jpg'.format(ref_id)))287 face_img_mask = cv2.imread(os.path.join(mask_dir, '{:d}.png'.format(ref_id)))288 289 rigid_mask = face_img_mask[..., 0] > 250290 rigid_num = np.sum(rigid_mask)291 flow_frame_num = 2500292 flow_frame_num = min(flow_frame_num, valid_img_num)293 rigid_flow = np.zeros((flow_frame_num, 2, rigid_num), np.float32)294 for i in range(flow_frame_num):295 flow = np.load(os.path.join(flow_dir, '{:d}_{:d}.npy'.format(ref_id, valid_img_ids[i])))296 rigid_flow[i] = flow[:, rigid_mask]297 rigid_flow = rigid_flow.transpose((2, 1, 0))298 rigid_flow = torch.as_tensor(rigid_flow).cuda()299 lap_kernel = torch.Tensor(300 (-0.5, 1.0, -0.5)).unsqueeze(0).unsqueeze(0).float().cuda()301 flow_lap = F.conv1d(302 rigid_flow.reshape(-1, 1, rigid_flow.shape[-1]), lap_kernel)303 flow_lap = flow_lap.view(rigid_flow.shape[0], 2, -1)304 flow_lap = torch.norm(flow_lap, dim=1)305 valid_frame = torch.mean(flow_lap, dim=0) < (torch.mean(flow_lap) * 3)306 flow_lap = flow_lap[:, valid_frame]307 rigid_flow_mean = torch.mean(flow_lap, dim=1)308 rigid_flow_show = (rigid_flow_mean - torch.min(rigid_flow_mean)) / \309 (torch.max(rigid_flow_mean) - torch.min(rigid_flow_mean)) * 255310 rigid_flow_show = rigid_flow_show.byte().cpu().numpy()311 rigid_flow_img = np.zeros((h, w, 1), dtype=np.uint8)312 rigid_flow_img[...] = 255313 rigid_flow_img[rigid_mask, 0] = rigid_flow_show314 cv2.imwrite(os.path.join(base_dir, 'rigid_flow.jpg'), rigid_flow_img)315 win_size, d_size = 5, 5316 sel_xys = np.zeros((h, w), dtype=np.int32)317 xys = []318 for y in range(0, h - win_size, win_size):319 for x in range(0, w - win_size, win_size):320 min_v = int(40)321 id_x = -1322 id_y = -1323 for dy in range(0, win_size):324 for dx in range(0, win_size):325 if rigid_flow_img[y + dy, x + dx, 0] < min_v:326 min_v = rigid_flow_img[y + dy, x + dx, 0]327 id_x = x + dx328 id_y = y + dy329 if id_x >= 0:330 if (np.sum(sel_xys[id_y - d_size:id_y + d_size + 1, id_x - d_size:id_x + d_size + 1]) == 0):331 cv2.circle(face_img, (id_x, id_y), 1, (255, 0, 0))332 xys.append(np.array((id_x, id_y), np.int32))333 sel_xys[id_y, id_x] = 1334 335 cv2.imwrite(os.path.join(base_dir, 'keypts.jpg'), face_img)336 np.savetxt(os.path.join(base_dir, 'keypoints.txt'), xys, '%d')337 key_xys = np.loadtxt(os.path.join(base_dir, 'keypoints.txt'), np.int32)338 track_xys = np.zeros((valid_img_num, key_xys.shape[0], 2), dtype=np.float32)339 track_dir = os.path.join(base_dir,'flow_result')340 track_paths = sorted(glob.glob(os.path.join(track_dir, '*.npy')), key=lambda x: int(x.split('/')[-1].split('.')[0]))341 342 for i, path in enumerate(track_paths):343 344 flow = np.load(path)345 for j in range(key_xys.shape[0]):346 x = key_xys[j, 0]347 y = key_xys[j, 1]348 track_xys[i, j, 0] = x + flow[0, y, x]349 track_xys[i, j, 1] = y + flow[1, y, x]350 np.save(os.path.join(base_dir, 'track_xys.npy'), track_xys)351 352 pose_opt_cmd = 'python data_utils/face_tracking/bundle_adjustment.py --path=' + base_dir + ' --img_h=' + \353 str(h) + ' --img_w=' + str(w)354 os.system(pose_opt_cmd)355 356def extract_blendshape(base_dir):357 print(f'[INFO] ===== extract blendshape =====')358 blendshape_cmd = 'python data_utils/blendshape_capture/main.py --path=' + base_dir359 os.system(blendshape_cmd)360 361 362def save_transforms(base_dir, ori_imgs_dir):363 print(f'[INFO] ===== save transforms =====')364 365 image_paths = glob.glob(os.path.join(ori_imgs_dir, '*.jpg'))366 367 # read one image to get H/W368 tmp_image = cv2.imread(image_paths[0], cv2.IMREAD_UNCHANGED) # [H, W, 3]369 h, w = tmp_image.shape[:2]370 371 params_dict = torch.load(os.path.join(base_dir, 'bundle_adjustment.pt'))372 focal_len = params_dict['focal']373 euler_angle = params_dict['euler']374 trans = params_dict['trans']375 valid_num = euler_angle.shape[0]376 377 train_val_split = int(valid_num * 10 / 11)378 train_ids = torch.arange(0, train_val_split)379 val_ids = torch.arange(train_val_split, valid_num)380 381 rot = euler2rot(euler_angle)382 rot_inv = rot.permute(0, 2, 1)383 trans_inv = -torch.bmm(rot_inv, trans.unsqueeze(2))384 385 pose = torch.eye(4, dtype=torch.float32)386 save_ids = ['train', 'val']387 train_val_ids = [train_ids, val_ids]388 mean_z = -float(torch.mean(trans[:, 2]).item())389 390 for split in range(2):391 transform_dict = dict()392 transform_dict['focal_len'] = float(focal_len[0])393 transform_dict['cx'] = float(w/2.0)394 transform_dict['cy'] = float(h/2.0)395 transform_dict['frames'] = []396 ids = train_val_ids[split]397 save_id = save_ids[split]398 399 for i in ids:400 i = i.item()401 frame_dict = dict()402 frame_dict['img_id'] = i403 frame_dict['aud_id'] = i404 405 pose[:3, :3] = rot_inv[i]406 pose[:3, 3] = trans_inv[i, :, 0]407 408 frame_dict['transform_matrix'] = pose.numpy().tolist()409 410 transform_dict['frames'].append(frame_dict)411 412 with open(os.path.join(base_dir, 'transforms_' + save_id + '.json'), 'w') as fp:413 json.dump(transform_dict, fp, indent=2, separators=(',', ': '))414 415 print(f'[INFO] ===== finished saving transforms =====')416 417 418if __name__ == '__main__':419 parser = argparse.ArgumentParser()420 parser.add_argument('path', type=str, help="path to video file")421 parser.add_argument('--task', type=int, default=-1, help="-1 means all")422 parser.add_argument('--asr', type=str, default='ave', help="ave, hubert or deepspeech")423 424 425 opt = parser.parse_args()426 427 base_dir = os.path.dirname(opt.path)428 429 wav_path = os.path.join(base_dir, 'aud.wav')430 ori_imgs_dir = os.path.join(base_dir, 'ori_imgs')431 parsing_dir = os.path.join(base_dir, 'parsing')432 gt_imgs_dir = os.path.join(base_dir, 'gt_imgs')433 torso_imgs_dir = os.path.join(base_dir, 'torso_imgs')434 mask_imgs_dir = os.path.join(base_dir, 'face_mask')435 flow_dir = os.path.join(base_dir, 'flow_result')436 437 438 os.makedirs(ori_imgs_dir, exist_ok=True)439 os.makedirs(parsing_dir, exist_ok=True)440 os.makedirs(gt_imgs_dir, exist_ok=True)441 os.makedirs(torso_imgs_dir, exist_ok=True)442 os.makedirs(mask_imgs_dir, exist_ok=True)443 os.makedirs(flow_dir, exist_ok=True)444 445 446 # extract audio447 if opt.task == -1 or opt.task == 1:448 extract_audio(opt.path, wav_path)449 extract_audio_features(wav_path, mode=opt.asr)450 451 # extract images452 if opt.task == -1 or opt.task == 2:453 extract_images(opt.path, ori_imgs_dir)454 455 # face parsing456 if opt.task == -1 or opt.task == 3:457 extract_semantics(ori_imgs_dir, parsing_dir)458 459 # extract bg460 if opt.task == -1 or opt.task == 4:461 extract_background(base_dir, ori_imgs_dir)462 463 # extract torso images and gt_images464 if opt.task == -1 or opt.task == 5:465 extract_torso_and_gt(base_dir, ori_imgs_dir)466 467 # extract face landmarks468 if opt.task == -1 or opt.task == 6:469 extract_landmarks(ori_imgs_dir)470 471 # face tracking472 if opt.task == -1 or opt.task == 7:473 face_tracking(ori_imgs_dir)474 475 # extract flow & pose optimization476 if opt.task == -1 or opt.task == 8:477 extract_flow(base_dir, ori_imgs_dir, mask_imgs_dir, flow_dir)478 479 # extract blendshape480 if opt.task == -1 or opt.task == 9:481 extract_blendshape(base_dir)482 483 # save transforms.json484 if opt.task == -1 or opt.task == 10:485 save_transforms(base_dir, ori_imgs_dir)486 487 