CoolFace
Apppublic

cnywt/SyncTalk

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
process.py487 linesDownload Raw Back to data_utils
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