CoolFace
Apppublic

RabbitRUI/ruispace

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
networks.py140 linesDownload Raw Back to audio2pose_models
1import torch.nn as nn2import torch3 4 5class ResidualConv(nn.Module):6    def __init__(self, input_dim, output_dim, stride, padding):7        super(ResidualConv, self).__init__()8 9        self.conv_block = nn.Sequential(10            nn.BatchNorm2d(input_dim),11            nn.ReLU(),12            nn.Conv2d(13                input_dim, output_dim, kernel_size=3, stride=stride, padding=padding14            ),15            nn.BatchNorm2d(output_dim),16            nn.ReLU(),17            nn.Conv2d(output_dim, output_dim, kernel_size=3, padding=1),18        )19        self.conv_skip = nn.Sequential(20            nn.Conv2d(input_dim, output_dim, kernel_size=3, stride=stride, padding=1),21            nn.BatchNorm2d(output_dim),22        )23 24    def forward(self, x):25 26        return self.conv_block(x) + self.conv_skip(x)27 28 29class Upsample(nn.Module):30    def __init__(self, input_dim, output_dim, kernel, stride):31        super(Upsample, self).__init__()32 33        self.upsample = nn.ConvTranspose2d(34            input_dim, output_dim, kernel_size=kernel, stride=stride35        )36 37    def forward(self, x):38        return self.upsample(x)39 40 41class Squeeze_Excite_Block(nn.Module):42    def __init__(self, channel, reduction=16):43        super(Squeeze_Excite_Block, self).__init__()44        self.avg_pool = nn.AdaptiveAvgPool2d(1)45        self.fc = nn.Sequential(46            nn.Linear(channel, channel // reduction, bias=False),47            nn.ReLU(inplace=True),48            nn.Linear(channel // reduction, channel, bias=False),49            nn.Sigmoid(),50        )51 52    def forward(self, x):53        b, c, _, _ = x.size()54        y = self.avg_pool(x).view(b, c)55        y = self.fc(y).view(b, c, 1, 1)56        return x * y.expand_as(x)57 58 59class ASPP(nn.Module):60    def __init__(self, in_dims, out_dims, rate=[6, 12, 18]):61        super(ASPP, self).__init__()62 63        self.aspp_block1 = nn.Sequential(64            nn.Conv2d(65                in_dims, out_dims, 3, stride=1, padding=rate[0], dilation=rate[0]66            ),67            nn.ReLU(inplace=True),68            nn.BatchNorm2d(out_dims),69        )70        self.aspp_block2 = nn.Sequential(71            nn.Conv2d(72                in_dims, out_dims, 3, stride=1, padding=rate[1], dilation=rate[1]73            ),74            nn.ReLU(inplace=True),75            nn.BatchNorm2d(out_dims),76        )77        self.aspp_block3 = nn.Sequential(78            nn.Conv2d(79                in_dims, out_dims, 3, stride=1, padding=rate[2], dilation=rate[2]80            ),81            nn.ReLU(inplace=True),82            nn.BatchNorm2d(out_dims),83        )84 85        self.output = nn.Conv2d(len(rate) * out_dims, out_dims, 1)86        self._init_weights()87 88    def forward(self, x):89        x1 = self.aspp_block1(x)90        x2 = self.aspp_block2(x)91        x3 = self.aspp_block3(x)92        out = torch.cat([x1, x2, x3], dim=1)93        return self.output(out)94 95    def _init_weights(self):96        for m in self.modules():97            if isinstance(m, nn.Conv2d):98                nn.init.kaiming_normal_(m.weight)99            elif isinstance(m, nn.BatchNorm2d):100                m.weight.data.fill_(1)101                m.bias.data.zero_()102 103 104class Upsample_(nn.Module):105    def __init__(self, scale=2):106        super(Upsample_, self).__init__()107 108        self.upsample = nn.Upsample(mode="bilinear", scale_factor=scale)109 110    def forward(self, x):111        return self.upsample(x)112 113 114class AttentionBlock(nn.Module):115    def __init__(self, input_encoder, input_decoder, output_dim):116        super(AttentionBlock, self).__init__()117 118        self.conv_encoder = nn.Sequential(119            nn.BatchNorm2d(input_encoder),120            nn.ReLU(),121            nn.Conv2d(input_encoder, output_dim, 3, padding=1),122            nn.MaxPool2d(2, 2),123        )124 125        self.conv_decoder = nn.Sequential(126            nn.BatchNorm2d(input_decoder),127            nn.ReLU(),128            nn.Conv2d(input_decoder, output_dim, 3, padding=1),129        )130 131        self.conv_attn = nn.Sequential(132            nn.BatchNorm2d(output_dim),133            nn.ReLU(),134            nn.Conv2d(output_dim, 1, 1),135        )136 137    def forward(self, x1, x2):138        out = self.conv_encoder(x1) + self.conv_decoder(x2)139        out = self.conv_attn(out)140        return out * x2