CoolFace
Apppublic

Shellbrady/LivePortrait5

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
helper.py155 linesDownload Raw Back to utils
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