| 923 | return prior_cam, p4_s_out, p3_s_out, p2_s_out, p1_s_out, p4_e_out, p3_e_out, p2_e_out, p1_e_out |
| 924 | |
| 925 | class REM_decoder_noSpade_noEdge(nn.Module): |
| 926 | def __init__(self, in_channels): |
| 927 | super(REM_decoder_noSpade_noEdge, self).__init__() |
| 928 | self.decoder4 = Decoder4_noEdge(in_channels)#可以弄4个decoder4 |
| 929 | self.decoder3 = Decoder4_noEdge(in_channels) |
| 930 | self.decoder2 = Decoder4_noEdge(in_channels) |
| 931 | self.decoder1 = Decoder4_noEdge(in_channels) |
| 932 | |
| 933 | def forward(self, x, prior_cam, pic): |
| 934 | f1, f2, f3, f4 = x |
| 935 | # 把所有的edge去掉。 |
| 936 | p4_s = self.decoder4(f4, prior_cam) |
| 937 | p4_s_out = F.interpolate(p4_s, size=pic.size()[2:], mode='bilinear') |
| 938 | |
| 939 | p3_s = self.decoder3(f3, p4_s) |
| 940 | p3_s_out = F.interpolate(p3_s, size=pic.size()[2:], mode='bilinear') |
| 941 | |
| 942 | p2_s = self.decoder2(f2, p3_s) |
| 943 | p2_s_out = F.interpolate(p2_s, size=pic.size()[2:], mode='bilinear') |
| 944 | # p2_e_out = F.interpolate(p2_e, size=pic.size()[2:], mode='bilinear') |
| 945 | |
| 946 | p1_s = self.decoder1(f1, p2_s) |
| 947 | p1_s_out = F.interpolate(p1_s, size=pic.size()[2:], mode='bilinear') |
| 948 | # p1_e_out = F.interpolate(p1_e, size=pic.size()[2:], mode='bilinear') |
| 949 | |
| 950 | return prior_cam, p4_s_out, p3_s_out, p2_s_out, p1_s_out |
| 951 | |
| 952 | class Decoder4_noRCAB(nn.Module): |
| 953 | def __init__(self, in_channels): |