Team Ai
Apppublic

Shellbrady/LivePortrait5

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
motion_extractor.py36 linesDownload Raw Back to modules
1# coding: utf-82 3"""4Motion extractor(M), which directly predicts the canonical keypoints, head pose and expression deformation of the input image5"""6 7from torch import nn8import torch9 10from .convnextv2 import convnextv2_tiny11from .util import filter_state_dict12 13model_dict = {14    'convnextv2_tiny': convnextv2_tiny,15}16 17 18class MotionExtractor(nn.Module):19    def __init__(self, **kwargs):20        super(MotionExtractor, self).__init__()21 22        # default is convnextv2_base23        backbone = kwargs.get('backbone', 'convnextv2_tiny')24        self.detector = model_dict.get(backbone)(**kwargs)25 26    def load_pretrained(self, init_path: str):27        if init_path not in (None, ''):28            state_dict = torch.load(init_path, map_location=lambda storage, loc: storage)['model']29            state_dict = filter_state_dict(state_dict, remove_name='head')30            ret = self.detector.load_state_dict(state_dict, strict=False)31            print(f'Load pretrained model from {init_path}, ret: {ret}')32 33    def forward(self, x):34        out = self.detector(x)35        return out36