| 972 | return f4_s, f4_e, p4_s, p4_e |
| 973 | |
| 974 | class Decoder_noRCAB(nn.Module): |
| 975 | def __init__(self, in_channels): |
| 976 | super(Decoder_noRCAB, self).__init__() |
| 977 | |
| 978 | self.spade_a = Spade(in_channels, in_channels) |
| 979 | self.spade_b = Spade(in_channels, in_channels) |
| 980 | |
| 981 | self.crb = BasicConv2d(in_channels*2, in_channels, kernel_size=3, padding=1) |
| 982 | self.out_e = nn.Conv2d(in_channels, 1, kernel_size=3, padding=1) |
| 983 | |
| 984 | self.conv3x3 = BasicConv2d(in_channels * 2, in_channels, kernel_size=3, padding=1) |
| 985 | self.out_s = nn.Conv2d(in_channels, 1, kernel_size=3, padding=1) |
| 986 | |
| 987 | def forward(self, f, f_s, f_e, p_s, p_e): |
| 988 | prior_cam = F.interpolate(p_s, size=f.size()[2:], mode='bilinear', align_corners=True) # 2,1,12,12->2,1,48,48 |
| 989 | r_prior_cam = 1 - torch.sigmoid(prior_cam) |
| 990 | prior_cam = torch.sigmoid(prior_cam) |
| 991 | |
| 992 | f_a = f * r_prior_cam.expand(-1, f.size()[1], -1, -1) + f |
| 993 | f_b = f * prior_cam.expand(-1, f.size()[1], -1, -1) + f |
| 994 | f_a = self.spade_a(f_a, f_e) |
| 995 | f_b = self.spade_b(f_b, f_e) |
| 996 | |
| 997 | f_s_new = self.conv3x3(torch.cat([f_a, f_b], 1)) |
| 998 | p_s_new = self.out_s(f_s_new) |
| 999 | |
| 1000 | p_e = F.interpolate(p_e, size=f.size()[2:], mode='bilinear', align_corners=True) |
| 1001 | f_s = F.interpolate(f_s, size=f.size()[2:], mode='bilinear', align_corners=True) |
| 1002 | |
| 1003 | f_e_new = self.crb(torch.cat([(f * p_e.expand(-1, f.size()[1], -1, -1) + f), f_s], 1)) |
| 1004 | p_e_new = self.out_e(f_e_new) |
| 1005 | |
| 1006 | return f_s_new, f_e_new, p_s_new, p_e_new |
| 1007 | |
| 1008 | class REM_decoder_noRCAB(nn.Module): |
| 1009 | def __init__(self, in_channels): |