| 791 | |
| 792 | |
| 793 | class REM_decoder(nn.Module): |
| 794 | def __init__(self, in_channels): |
| 795 | super(REM_decoder, self).__init__() |
| 796 | self.decoder4 = Decoder4(in_channels) |
| 797 | self.decoder3 = Decoder(in_channels) |
| 798 | self.decoder2 = Decoder(in_channels) |
| 799 | self.decoder1 = Decoder(in_channels) |
| 800 | |
| 801 | def forward(self, x, prior_cam, pic): |
| 802 | f1, f2, f3, f4 = x |
| 803 | f4_s, f4_e, p4_s, p4_e = self.decoder4(f4, prior_cam) |
| 804 | p4_s_out = F.interpolate(p4_s, size=pic.size()[2:], mode='bilinear') |
| 805 | p4_e_out = F.interpolate(p4_e, size=pic.size()[2:], mode='bilinear') |
| 806 | |
| 807 | f3_s, f3_e, p3_s, p3_e = self.decoder3(f3, f4_s, f4_e, p4_s, p4_e) |
| 808 | p3_s_out = F.interpolate(p3_s, size=pic.size()[2:], mode='bilinear') |
| 809 | p3_e_out = F.interpolate(p3_e, size=pic.size()[2:], mode='bilinear') |
| 810 | |
| 811 | f2_s, f2_e, p2_s, p2_e = self.decoder2(f2, f3_s, f3_e, p3_s, p3_e) |
| 812 | p2_s_out = F.interpolate(p2_s, size=pic.size()[2:], mode='bilinear') |
| 813 | p2_e_out = F.interpolate(p2_e, size=pic.size()[2:], mode='bilinear') |
| 814 | |
| 815 | f1_s, f1_e, p1_s, p1_e = self.decoder1(f1, f2_s, f2_e, p2_s, p2_e) |
| 816 | p1_s_out = F.interpolate(p1_s, size=pic.size()[2:], mode='bilinear') |
| 817 | p1_e_out = F.interpolate(p1_e, size=pic.size()[2:], mode='bilinear') |
| 818 | |
| 819 | 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 |
| 820 | |
| 821 | |
| 822 | """ |
nothing calls this directly
no outgoing calls
no test coverage detected