| 708 | |
| 709 | |
| 710 | class Decoder4(nn.Module): |
| 711 | def __init__(self, in_channels): |
| 712 | super(Decoder4, self).__init__() |
| 713 | self.rcab_a = RCAB(in_channels) |
| 714 | self.rcab_b = RCAB(in_channels) |
| 715 | self.crb = BasicConv2d(in_channels, in_channels, kernel_size=3, padding=1) |
| 716 | self.out_e = nn.Conv2d(in_channels, 1, kernel_size=3, padding=1) |
| 717 | |
| 718 | self.conv3x3 = BasicConv2d(in_channels * 2, in_channels, kernel_size=3, padding=1) |
| 719 | self.out_s = nn.Conv2d(in_channels, 1, kernel_size=3, padding=1) |
| 720 | |
| 721 | def forward(self, f4, prior_cam): |
| 722 | prior_cam = F.interpolate(prior_cam, size=f4.size()[2:], mode='bilinear', align_corners=True) # 2,1,12,12->2,1,48,48 |
| 723 | r_prior_cam = 1 - torch.sigmoid(prior_cam) |
| 724 | prior_cam = torch.sigmoid(prior_cam) |
| 725 | f4_a = self.rcab_a(f4 * r_prior_cam.expand(-1, f4.size()[1], -1, -1) + f4) |
| 726 | f4_b = self.rcab_b(f4 * prior_cam.expand(-1, f4.size()[1], -1, -1) + f4) |
| 727 | f4_s = self.conv3x3(torch.cat([f4_a, f4_b], 1)) |
| 728 | p4_s = self.out_s(f4_s) |
| 729 | f4_e = self.crb(f4) |
| 730 | p4_e = self.out_e(f4_e) |
| 731 | |
| 732 | return f4_s, f4_e, p4_s, p4_e |
| 733 | |
| 734 | |
| 735 | class Spade(nn.Module): |