CoolFace
Apppublic

HexB/CodeFormer

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
bisenet.py141 linesDownload Raw Back to parsing
1import torch2import torch.nn as nn3import torch.nn.functional as F4 5from .resnet import ResNet186 7 8class ConvBNReLU(nn.Module):9 10    def __init__(self, in_chan, out_chan, ks=3, stride=1, padding=1):11        super(ConvBNReLU, self).__init__()12        self.conv = nn.Conv2d(in_chan, out_chan, kernel_size=ks, stride=stride, padding=padding, bias=False)13        self.bn = nn.BatchNorm2d(out_chan)14 15    def forward(self, x):16        x = self.conv(x)17        x = F.relu(self.bn(x))18        return x19 20 21class BiSeNetOutput(nn.Module):22 23    def __init__(self, in_chan, mid_chan, num_class):24        super(BiSeNetOutput, self).__init__()25        self.conv = ConvBNReLU(in_chan, mid_chan, ks=3, stride=1, padding=1)26        self.conv_out = nn.Conv2d(mid_chan, num_class, kernel_size=1, bias=False)27 28    def forward(self, x):29        feat = self.conv(x)30        out = self.conv_out(feat)31        return out, feat32 33 34class AttentionRefinementModule(nn.Module):35 36    def __init__(self, in_chan, out_chan):37        super(AttentionRefinementModule, self).__init__()38        self.conv = ConvBNReLU(in_chan, out_chan, ks=3, stride=1, padding=1)39        self.conv_atten = nn.Conv2d(out_chan, out_chan, kernel_size=1, bias=False)40        self.bn_atten = nn.BatchNorm2d(out_chan)41        self.sigmoid_atten = nn.Sigmoid()42 43    def forward(self, x):44        feat = self.conv(x)45        atten = F.avg_pool2d(feat, feat.size()[2:])46        atten = self.conv_atten(atten)47        atten = self.bn_atten(atten)48        atten = self.sigmoid_atten(atten)49        out = torch.mul(feat, atten)50        return out51 52 53class ContextPath(nn.Module):54 55    def __init__(self):56        super(ContextPath, self).__init__()57        self.resnet = ResNet18()58        self.arm16 = AttentionRefinementModule(256, 128)59        self.arm32 = AttentionRefinementModule(512, 128)60        self.conv_head32 = ConvBNReLU(128, 128, ks=3, stride=1, padding=1)61        self.conv_head16 = ConvBNReLU(128, 128, ks=3, stride=1, padding=1)62        self.conv_avg = ConvBNReLU(512, 128, ks=1, stride=1, padding=0)63 64    def forward(self, x):65        feat8, feat16, feat32 = self.resnet(x)66        h8, w8 = feat8.size()[2:]67        h16, w16 = feat16.size()[2:]68        h32, w32 = feat32.size()[2:]69 70        avg = F.avg_pool2d(feat32, feat32.size()[2:])71        avg = self.conv_avg(avg)72        avg_up = F.interpolate(avg, (h32, w32), mode='nearest')73 74        feat32_arm = self.arm32(feat32)75        feat32_sum = feat32_arm + avg_up76        feat32_up = F.interpolate(feat32_sum, (h16, w16), mode='nearest')77        feat32_up = self.conv_head32(feat32_up)78 79        feat16_arm = self.arm16(feat16)80        feat16_sum = feat16_arm + feat32_up81        feat16_up = F.interpolate(feat16_sum, (h8, w8), mode='nearest')82        feat16_up = self.conv_head16(feat16_up)83 84        return feat8, feat16_up, feat32_up  # x8, x8, x1685 86 87class FeatureFusionModule(nn.Module):88 89    def __init__(self, in_chan, out_chan):90        super(FeatureFusionModule, self).__init__()91        self.convblk = ConvBNReLU(in_chan, out_chan, ks=1, stride=1, padding=0)92        self.conv1 = nn.Conv2d(out_chan, out_chan // 4, kernel_size=1, stride=1, padding=0, bias=False)93        self.conv2 = nn.Conv2d(out_chan // 4, out_chan, kernel_size=1, stride=1, padding=0, bias=False)94        self.relu = nn.ReLU(inplace=True)95        self.sigmoid = nn.Sigmoid()96 97    def forward(self, fsp, fcp):98        fcat = torch.cat([fsp, fcp], dim=1)99        feat = self.convblk(fcat)100        atten = F.avg_pool2d(feat, feat.size()[2:])101        atten = self.conv1(atten)102        atten = self.relu(atten)103        atten = self.conv2(atten)104        atten = self.sigmoid(atten)105        feat_atten = torch.mul(feat, atten)106        feat_out = feat_atten + feat107        return feat_out108 109 110class BiSeNet(nn.Module):111 112    def __init__(self, num_class):113        super(BiSeNet, self).__init__()114        self.cp = ContextPath()115        self.ffm = FeatureFusionModule(256, 256)116        self.conv_out = BiSeNetOutput(256, 256, num_class)117        self.conv_out16 = BiSeNetOutput(128, 64, num_class)118        self.conv_out32 = BiSeNetOutput(128, 64, num_class)119 120    def forward(self, x, return_feat=False):121        h, w = x.size()[2:]122        feat_res8, feat_cp8, feat_cp16 = self.cp(x)  # return res3b1 feature123        feat_sp = feat_res8  # replace spatial path feature with res3b1 feature124        feat_fuse = self.ffm(feat_sp, feat_cp8)125 126        out, feat = self.conv_out(feat_fuse)127        out16, feat16 = self.conv_out16(feat_cp8)128        out32, feat32 = self.conv_out32(feat_cp16)129 130        out = F.interpolate(out, (h, w), mode='bilinear', align_corners=True)131        out16 = F.interpolate(out16, (h, w), mode='bilinear', align_corners=True)132        out32 = F.interpolate(out32, (h, w), mode='bilinear', align_corners=True)133 134        if return_feat:135            feat = F.interpolate(feat, (h, w), mode='bilinear', align_corners=True)136            feat16 = F.interpolate(feat16, (h, w), mode='bilinear', align_corners=True)137            feat32 = F.interpolate(feat32, (h, w), mode='bilinear', align_corners=True)138            return out, out16, out32, feat, feat16, feat32139        else:140            return out, out16, out32141