MCPcopy Create free account
hub / github.com/ChunmingHe/WS-SAM / Decoder

Class Decoder

lib/VIT/decoder/decoder_p.py:220–294  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

218
219
220class Decoder(nn.Module):
221 def __init__(self, channels):
222 super(Decoder, self).__init__()
223
224 self.side_conv1 = nn.Conv2d(512, channels, kernel_size=3, stride=1, padding=1)
225 self.side_conv2 = nn.Conv2d(320, channels, kernel_size=3, stride=1, padding=1)
226 self.side_conv3 = nn.Conv2d(128, channels, kernel_size=3, stride=1, padding=1)
227 self.side_conv4 = nn.Conv2d(64, channels, kernel_size=3, stride=1, padding=1)
228
229 self.conv_block = Conv_Block(channels)
230
231 self.fuse1 = nn.Sequential(nn.Conv2d(channels*2, channels, kernel_size=3, stride=1, padding=1, bias=False),nn.BatchNorm2d(channels))
232 self.fuse2 = nn.Sequential(nn.Conv2d(channels*2, channels, kernel_size=3, stride=1, padding=1, bias=False),nn.BatchNorm2d(channels))
233 self.fuse3 = nn.Sequential(nn.Conv2d(channels*2, channels, kernel_size=3, stride=1, padding=1, bias=False),nn.BatchNorm2d(channels))
234
235 self.MSA5=MSA_module(dim = channels)
236 self.MSA4=MSA_module(dim = channels)
237 self.MSA3=MSA_module(dim = channels)
238 self.MSA2=MSA_module(dim = channels)
239
240 self.predtrans1 = nn.Conv2d(channels, 1, kernel_size=3, padding=1)
241 self.predtrans2 = nn.Conv2d(channels, 1, kernel_size=3, padding=1)
242 self.predtrans3 = nn.Conv2d(channels, 1, kernel_size=3, padding=1)
243 self.predtrans4 = nn.Conv2d(channels, 1, kernel_size=3, padding=1)
244 self.predtrans5 = nn.Conv2d(channels, 1, kernel_size=3, padding=1)
245
246 self.initialize()
247
248
249
250 def forward(self, E4, E3, E2, E1,shape):
251 E4, E3, E2, E1= self.side_conv1(E4), self.side_conv2(E3), self.side_conv3(E2), self.side_conv4(E1)
252
253 if E4.size()[2:] != E3.size()[2:]:
254 E4 = F.interpolate(E4, size=E3.size()[2:], mode='bilinear')
255 if E2.size()[2:] != E3.size()[2:]:
256 E2 = F.interpolate(E2, size=E3.size()[2:], mode='bilinear')
257
258 E5 = self.conv_block(E4, E3, E2)
259
260 E4 = torch.cat((E4, E5),1)
261 E3 = torch.cat((E3, E5),1)
262 E2 = torch.cat((E2, E5),1)
263
264 E4 = F.relu(self.fuse1(E4), inplace=True)
265 E3 = F.relu(self.fuse2(E3), inplace=True)
266 E2 = F.relu(self.fuse3(E2), inplace=True)
267
268 P5 = self.predtrans5(E5)
269
270 D4 = self.MSA5(E5, E4, P5)
271 D4 = F.interpolate(D4, size=E3.size()[2:], mode='bilinear')
272 P4 = self.predtrans4(D4)
273
274 D3 = self.MSA4(D4, E3, P4)
275 D3 = F.interpolate(D3, size=E2.size()[2:], mode='bilinear')
276 P3 = self.predtrans3(D3)
277

Callers 1

__init__Method · 0.50

Calls

no outgoing calls

Tested by

no test coverage detected