MCPcopy Create free account
hub / github.com/DIVE128/DMVSNet / CostRegNet_part_refine

Class CostRegNet_part_refine

networks/module.py:400–436  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

398 return x
399
400class CostRegNet_part_refine(nn.Module):
401 def __init__(self, in_channels, base_channels,stage=0):
402 super(CostRegNet_part_refine, self).__init__()
403 self.conv0 = Conv3d(in_channels, base_channels, padding=1)
404
405 self.conv1 = Conv3d(base_channels, base_channels * 2, stride=2, padding=1)
406 self.conv2 = Conv3d(base_channels * 2, base_channels * 2, padding=1)
407
408 self.conv3 = Conv3d(base_channels * 2, base_channels * 4, stride=2, padding=1)
409 self.conv4 = Conv3d(base_channels * 4, base_channels * 4, padding=1)
410
411 self.conv5 = Conv2d(base_channels * 4, base_channels * 8,3, stride=2, padding=1)
412 self.conv6 = Conv2d(base_channels * 8, base_channels * 8,3, padding=1)
413
414 self.conv7 = Deconv2d(base_channels * 8, base_channels * 4,3, stride=2, padding=1, output_padding=1)
415
416 self.conv9 = Deconv3d(base_channels * 4, base_channels * 2, stride=2, padding=1, output_padding=1)
417
418 self.conv11 = Deconv3d(base_channels * 2, base_channels * 1, stride=2, padding=1, output_padding=1)
419
420 # self.prob = nn.Conv3d(base_channels, 1 if stage==0 else 2, 3, stride=1, padding=1, bias=False)
421 self.prob = nn.Conv3d(base_channels, 2, 3, stride=1, padding=1, bias=False)
422
423
424
425
426 def forward(self, x,stage=0):
427 conv0 = self.conv0(x)
428 conv2 = self.conv2(self.conv1(conv0))
429 conv4 = self.conv4(self.conv3(conv2)).squeeze(2)
430 x=self.conv6(self.conv5(conv4))
431 x=conv4+self.conv7(x)
432 x=x.unsqueeze(2)
433 x = conv2 + self.conv9(x)
434 x = conv0 + self.conv11(x)
435 x = self.prob(x)
436 return x
437class AggWeightNetVolume(nn.Module):
438 def __init__(self, in_channels=32,hid_channels=1,out_channels=1,relu=True):
439 super(AggWeightNetVolume, self).__init__()

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected