Team Ai
Apppublic

Shellbrady/LivePortrait5

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
stitching_retargeting_network.py39 linesDownload Raw Back to modules
1# coding: utf-82 3"""4Stitching module(S) and two retargeting modules(R) defined in the paper.5 6- The stitching module pastes the animated portrait back into the original image space without pixel misalignment, such as in7the stitching region.8 9- The eyes retargeting module is designed to address the issue of incomplete eye closure during cross-id reenactment, especially10when a person with small eyes drives a person with larger eyes.11 12- The lip retargeting module is designed similarly to the eye retargeting module, and can also normalize the input by ensuring that13the lips are in a closed state, which facilitates better animation driving.14"""15from torch import nn16 17 18class StitchingRetargetingNetwork(nn.Module):19    def __init__(self, input_size, hidden_sizes, output_size):20        super(StitchingRetargetingNetwork, self).__init__()21        layers = []22        for i in range(len(hidden_sizes)):23            if i == 0:24                layers.append(nn.Linear(input_size, hidden_sizes[i]))25            else:26                layers.append(nn.Linear(hidden_sizes[i - 1], hidden_sizes[i]))27            layers.append(nn.ReLU(inplace=True))28        layers.append(nn.Linear(hidden_sizes[-1], output_size))29        self.mlp = nn.Sequential(*layers)30 31    def initialize_weights_to_zero(self):32        for m in self.modules():33            if isinstance(m, nn.Linear):34                nn.init.zeros_(m.weight)35                nn.init.zeros_(m.bias)36 37    def forward(self, x):38        return self.mlp(x)39