CoolFace
Apppublic

Shellbrady/LivePortrait5

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
spade_generator.py59 linesDownload Raw Back to modules
1# coding: utf-82 3"""4Spade decoder(G) defined in the paper, which input the warped feature to generate the animated image.5"""6 7import torch8from torch import nn9import torch.nn.functional as F10from .util import SPADEResnetBlock11 12 13class SPADEDecoder(nn.Module):14    def __init__(self, upscale=1, max_features=256, block_expansion=64, out_channels=64, num_down_blocks=2):15        for i in range(num_down_blocks):16            input_channels = min(max_features, block_expansion * (2 ** (i + 1)))17        self.upscale = upscale18        super().__init__()19        norm_G = 'spadespectralinstance'20        label_num_channels = input_channels  # 25621 22        self.fc = nn.Conv2d(input_channels, 2 * input_channels, 3, padding=1)23        self.G_middle_0 = SPADEResnetBlock(2 * input_channels, 2 * input_channels, norm_G, label_num_channels)24        self.G_middle_1 = SPADEResnetBlock(2 * input_channels, 2 * input_channels, norm_G, label_num_channels)25        self.G_middle_2 = SPADEResnetBlock(2 * input_channels, 2 * input_channels, norm_G, label_num_channels)26        self.G_middle_3 = SPADEResnetBlock(2 * input_channels, 2 * input_channels, norm_G, label_num_channels)27        self.G_middle_4 = SPADEResnetBlock(2 * input_channels, 2 * input_channels, norm_G, label_num_channels)28        self.G_middle_5 = SPADEResnetBlock(2 * input_channels, 2 * input_channels, norm_G, label_num_channels)29        self.up_0 = SPADEResnetBlock(2 * input_channels, input_channels, norm_G, label_num_channels)30        self.up_1 = SPADEResnetBlock(input_channels, out_channels, norm_G, label_num_channels)31        self.up = nn.Upsample(scale_factor=2)32 33        if self.upscale is None or self.upscale <= 1:34            self.conv_img = nn.Conv2d(out_channels, 3, 3, padding=1)35        else:36            self.conv_img = nn.Sequential(37                nn.Conv2d(out_channels, 3 * (2 * 2), kernel_size=3, padding=1),38                nn.PixelShuffle(upscale_factor=2)39            )40 41    def forward(self, feature):42        seg = feature  # Bx256x64x6443        x = self.fc(feature)  # Bx512x64x6444        x = self.G_middle_0(x, seg)45        x = self.G_middle_1(x, seg)46        x = self.G_middle_2(x, seg)47        x = self.G_middle_3(x, seg)48        x = self.G_middle_4(x, seg)49        x = self.G_middle_5(x, seg)50 51        x = self.up(x)  # Bx512x64x64 -> Bx512x128x12852        x = self.up_0(x, seg)  # Bx512x128x128 -> Bx256x128x12853        x = self.up(x)  # Bx256x128x128 -> Bx256x256x25654        x = self.up_1(x, seg)  # Bx256x256x256 -> Bx64x256x25655 56        x = self.conv_img(F.leaky_relu(x, 2e-1))  # Bx64x256x256 -> Bx3xHxW57        x = torch.sigmoid(x)  # Bx3xHxW58 59        return x