| 150 | feat_align = self.relu(self.dcpack_L2([feat_up, offset])) |
| 151 | return feat_align, feat_arm |
| 152 | class BiseNetDecoder(nn.Module): |
| 153 | def __init__(self, num_classes, channels): |
| 154 | super(BiseNetDecoder, self).__init__() |
| 155 | channels8,channels16=channels["8"],channels["16"] |
| 156 | self.arm16 = AttentionRefinementModule(channels16, 128) |
| 157 | self.conv_head16 = ConvBnAct(128,128,3,1,1) |
| 158 | self.avg_pool=nn.AdaptiveAvgPool2d(1) |
| 159 | self.conv_avg = ConvBnAct(channels16,128) |
| 160 | self.ffm=FeatureFusionModule(128+channels8,128) |
| 161 | self.conv=ConvBnAct(128,128,3,1,1) |
| 162 | self.classifier=nn.Conv2d(128, num_classes, 1) |
| 163 | |
| 164 | def forward(self, x): |
| 165 | x8,x16= x["8"], x["16"] |
| 166 | |
| 167 | avg=self.avg_pool(x16) |
| 168 | avg = self.conv_avg(avg) |
| 169 | avg_up = F.interpolate(avg, size=x16.shape[-2:], mode='nearest') |
| 170 | |
| 171 | x16 = self.arm16(x16) |
| 172 | x16 = x16 + avg_up |
| 173 | x16 = F.interpolate(x16, size=x8.shape[-2:], mode='nearest') |
| 174 | x16 = self.conv_head16(x16) |
| 175 | |
| 176 | x=self.ffm(x8,x16) |
| 177 | x=self.conv(x) |
| 178 | x=self.classifier(x) |
| 179 | return x |
| 180 | class SFNetDecoder(nn.Module): |
| 181 | def __init__(self, num_classes, channels, fpn_dim=64, fpn_dsn=False): |
| 182 | super().__init__() |