| 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 | |