MCPcopy Create free account
hub / github.com/SLDGroup/EMCAD / forward

Method forward

lib/networks.py:88–115  ·  view source on GitHub ↗
(self, x, mode='test')

Source from the content-addressed store, hash-verified

86 self.out_head1 = nn.Conv2d(channels[3], num_classes, 1)
87
88 def forward(self, x, mode='test'):
89
90 # if grayscale input, convert to 3 channels
91 if x.size()[1] == 1:
92 x = self.conv(x)
93
94 # encoder
95 x1, x2, x3, x4 = self.backbone(x)
96 #print(x1.shape, x2.shape, x3.shape, x4.shape)
97
98 # decoder
99 dec_outs = self.decoder(x4, [x3, x2, x1])
100
101 # prediction heads
102 p4 = self.out_head4(dec_outs[0])
103 p3 = self.out_head3(dec_outs[1])
104 p2 = self.out_head2(dec_outs[2])
105 p1 = self.out_head1(dec_outs[3])
106
107 p4 = F.interpolate(p4, scale_factor=32, mode='bilinear')
108 p3 = F.interpolate(p3, scale_factor=16, mode='bilinear')
109 p2 = F.interpolate(p2, scale_factor=8, mode='bilinear')
110 p1 = F.interpolate(p1, scale_factor=4, mode='bilinear')
111
112 if mode == 'test':
113 return [p4, p3, p2, p1]
114
115 return [p4, p3, p2, p1]
116
117
118

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected