| 218 | |
| 219 | |
| 220 | class Decoder(nn.Module): |
| 221 | def __init__(self, channels): |
| 222 | super(Decoder, self).__init__() |
| 223 | |
| 224 | self.side_conv1 = nn.Conv2d(512, channels, kernel_size=3, stride=1, padding=1) |
| 225 | self.side_conv2 = nn.Conv2d(320, channels, kernel_size=3, stride=1, padding=1) |
| 226 | self.side_conv3 = nn.Conv2d(128, channels, kernel_size=3, stride=1, padding=1) |
| 227 | self.side_conv4 = nn.Conv2d(64, channels, kernel_size=3, stride=1, padding=1) |
| 228 | |
| 229 | self.conv_block = Conv_Block(channels) |
| 230 | |
| 231 | self.fuse1 = nn.Sequential(nn.Conv2d(channels*2, channels, kernel_size=3, stride=1, padding=1, bias=False),nn.BatchNorm2d(channels)) |
| 232 | self.fuse2 = nn.Sequential(nn.Conv2d(channels*2, channels, kernel_size=3, stride=1, padding=1, bias=False),nn.BatchNorm2d(channels)) |
| 233 | self.fuse3 = nn.Sequential(nn.Conv2d(channels*2, channels, kernel_size=3, stride=1, padding=1, bias=False),nn.BatchNorm2d(channels)) |
| 234 | |
| 235 | self.MSA5=MSA_module(dim = channels) |
| 236 | self.MSA4=MSA_module(dim = channels) |
| 237 | self.MSA3=MSA_module(dim = channels) |
| 238 | self.MSA2=MSA_module(dim = channels) |
| 239 | |
| 240 | self.predtrans1 = nn.Conv2d(channels, 1, kernel_size=3, padding=1) |
| 241 | self.predtrans2 = nn.Conv2d(channels, 1, kernel_size=3, padding=1) |
| 242 | self.predtrans3 = nn.Conv2d(channels, 1, kernel_size=3, padding=1) |
| 243 | self.predtrans4 = nn.Conv2d(channels, 1, kernel_size=3, padding=1) |
| 244 | self.predtrans5 = nn.Conv2d(channels, 1, kernel_size=3, padding=1) |
| 245 | |
| 246 | self.initialize() |
| 247 | |
| 248 | |
| 249 | |
| 250 | def forward(self, E4, E3, E2, E1,shape): |
| 251 | E4, E3, E2, E1= self.side_conv1(E4), self.side_conv2(E3), self.side_conv3(E2), self.side_conv4(E1) |
| 252 | |
| 253 | if E4.size()[2:] != E3.size()[2:]: |
| 254 | E4 = F.interpolate(E4, size=E3.size()[2:], mode='bilinear') |
| 255 | if E2.size()[2:] != E3.size()[2:]: |
| 256 | E2 = F.interpolate(E2, size=E3.size()[2:], mode='bilinear') |
| 257 | |
| 258 | E5 = self.conv_block(E4, E3, E2) |
| 259 | |
| 260 | E4 = torch.cat((E4, E5),1) |
| 261 | E3 = torch.cat((E3, E5),1) |
| 262 | E2 = torch.cat((E2, E5),1) |
| 263 | |
| 264 | E4 = F.relu(self.fuse1(E4), inplace=True) |
| 265 | E3 = F.relu(self.fuse2(E3), inplace=True) |
| 266 | E2 = F.relu(self.fuse3(E2), inplace=True) |
| 267 | |
| 268 | P5 = self.predtrans5(E5) |
| 269 | |
| 270 | D4 = self.MSA5(E5, E4, P5) |
| 271 | D4 = F.interpolate(D4, size=E3.size()[2:], mode='bilinear') |
| 272 | P4 = self.predtrans4(D4) |
| 273 | |
| 274 | D3 = self.MSA4(D4, E3, P4) |
| 275 | D3 = F.interpolate(D3, size=E2.size()[2:], mode='bilinear') |
| 276 | P3 = self.predtrans3(D3) |
| 277 | |