Shellbrady/LivePortrait5
0
1# coding: utf-82 3"""4utility functions and classes to handle feature extraction and model loading5"""6 7import os8import os.path as osp9import torch10from collections import OrderedDict11 12from ..modules.spade_generator import SPADEDecoder13from ..modules.warping_network import WarpingNetwork14from ..modules.motion_extractor import MotionExtractor15from ..modules.appearance_feature_extractor import AppearanceFeatureExtractor16from ..modules.stitching_retargeting_network import StitchingRetargetingNetwork17 18 19def suffix(filename):20 """a.jpg -> jpg"""21 pos = filename.rfind(".")22 if pos == -1:23 return ""24 return filename[pos + 1:]25 26 27def prefix(filename):28 """a.jpg -> a"""29 pos = filename.rfind(".")30 if pos == -1:31 return filename32 return filename[:pos]33 34 35def basename(filename):36 """a/b/c.jpg -> c"""37 return prefix(osp.basename(filename))38 39 40def is_video(file_path):41 if file_path.lower().endswith((".mp4", ".mov", ".avi", ".webm")) or osp.isdir(file_path):42 return True43 return False44 45 46def is_template(file_path):47 if file_path.endswith(".pkl"):48 return True49 return False50 51 52def mkdir(d, log=False):53 # return self-assined `d`, for one line code54 if not osp.exists(d):55 os.makedirs(d, exist_ok=True)56 if log:57 print(f"Make dir: {d}")58 return d59 60 61def squeeze_tensor_to_numpy(tensor):62 out = tensor.data.squeeze(0).cpu().numpy()63 return out64 65 66def dct2cuda(dct: dict, device_id: int):67 for key in dct:68 dct[key] = torch.tensor(dct[key]).cuda(device_id)69 return dct70 71 72def concat_feat(kp_source: torch.Tensor, kp_driving: torch.Tensor) -> torch.Tensor:73 """74 kp_source: (bs, k, 3)75 kp_driving: (bs, k, 3)76 Return: (bs, 2k*3)77 """78 bs_src = kp_source.shape[0]79 bs_dri = kp_driving.shape[0]80 assert bs_src == bs_dri, 'batch size must be equal'81 82 feat = torch.cat([kp_source.view(bs_src, -1), kp_driving.view(bs_dri, -1)], dim=1)83 return feat84 85 86def remove_ddp_dumplicate_key(state_dict):87 state_dict_new = OrderedDict()88 for key in state_dict.keys():89 state_dict_new[key.replace('module.', '')] = state_dict[key]90 return state_dict_new91 92 93def load_model(ckpt_path, model_config, device, model_type):94 model_params = model_config['model_params'][f'{model_type}_params']95 96 if model_type == 'appearance_feature_extractor':97 model = AppearanceFeatureExtractor(**model_params).cuda(device)98 elif model_type == 'motion_extractor':99 model = MotionExtractor(**model_params).cuda(device)100 elif model_type == 'warping_module':101 model = WarpingNetwork(**model_params).cuda(device)102 elif model_type == 'spade_generator':103 model = SPADEDecoder(**model_params).cuda(device)104 elif model_type == 'stitching_retargeting_module':105 # Special handling for stitching and retargeting module106 config = model_config['model_params']['stitching_retargeting_module_params']107 checkpoint = torch.load(ckpt_path, map_location=lambda storage, loc: storage)108 109 stitcher = StitchingRetargetingNetwork(**config.get('stitching'))110 stitcher.load_state_dict(remove_ddp_dumplicate_key(checkpoint['retarget_shoulder']))111 stitcher = stitcher.cuda(device)112 stitcher.eval()113 114 retargetor_lip = StitchingRetargetingNetwork(**config.get('lip'))115 retargetor_lip.load_state_dict(remove_ddp_dumplicate_key(checkpoint['retarget_mouth']))116 retargetor_lip = retargetor_lip.cuda(device)117 retargetor_lip.eval()118 119 retargetor_eye = StitchingRetargetingNetwork(**config.get('eye'))120 retargetor_eye.load_state_dict(remove_ddp_dumplicate_key(checkpoint['retarget_eye']))121 retargetor_eye = retargetor_eye.cuda(device)122 retargetor_eye.eval()123 124 return {125 'stitching': stitcher,126 'lip': retargetor_lip,127 'eye': retargetor_eye128 }129 else:130 raise ValueError(f"Unknown model type: {model_type}")131 132 model.load_state_dict(torch.load(ckpt_path, map_location=lambda storage, loc: storage))133 model.eval()134 return model135 136 137# get coefficients of Eqn. 7138def calculate_transformation(config, s_kp_info, t_0_kp_info, t_i_kp_info, R_s, R_t_0, R_t_i):139 if config.relative:140 new_rotation = (R_t_i @ R_t_0.permute(0, 2, 1)) @ R_s141 new_expression = s_kp_info['exp'] + (t_i_kp_info['exp'] - t_0_kp_info['exp'])142 else:143 new_rotation = R_t_i144 new_expression = t_i_kp_info['exp']145 new_translation = s_kp_info['t'] + (t_i_kp_info['t'] - t_0_kp_info['t'])146 new_translation[..., 2].fill_(0) # Keep the z-axis unchanged147 new_scale = s_kp_info['scale'] * (t_i_kp_info['scale'] / t_0_kp_info['scale'])148 return new_rotation, new_expression, new_translation, new_scale149 150 151def load_description(fp):152 with open(fp, 'r', encoding='utf-8') as f:153 content = f.read()154 return content155 