HexB/CodeFormer
0
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 