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

Class CostRegNet_part

networks/module.py:358–398  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

356 results=torch.cat((self.cosR_small(x),self.cosR_huge(x)),axis=1)
357 return results
358class CostRegNet_part(nn.Module):
359 def __init__(self, in_channels, base_channels,stage=0):
360 super(CostRegNet_part, self).__init__()
361 self.conv0 = Conv3d(in_channels, base_channels, padding=1)
362
363 self.conv1 = Conv3d(base_channels, base_channels * 2, stride=2, padding=1)
364 self.conv2 = Conv3d(base_channels * 2, base_channels * 2, padding=1)
365
366 self.conv3 = Conv3d(base_channels * 2, base_channels * 4, stride=2, padding=1)
367 self.conv4 = Conv3d(base_channels * 4, base_channels * 4, padding=1)
368
369 self.conv5 = Conv3d(base_channels * 4, base_channels * 8, stride=2, padding=1)
370 self.conv6 = Conv3d(base_channels * 8, base_channels * 8, padding=1)
371
372 self.conv7 = Deconv3d(base_channels * 8, base_channels * 4, stride=2, padding=1, output_padding=1)
373
374 self.conv9 = Deconv3d(base_channels * 4, base_channels * 2, stride=2, padding=1, output_padding=1)
375
376 self.conv11 = Deconv3d(base_channels * 2, base_channels * 1, stride=2, padding=1, output_padding=1)
377
378 # self.prob = nn.Conv3d(base_channels, 1 if stage==0 else 2, 3, stride=1, padding=1, bias=False)
379 self.prob = nn.Conv3d(base_channels, 2, 3, stride=1, padding=1, bias=False)
380
381
382 # for m in self.modules():
383 # if isinstance(m, nn.Conv2d):
384 # nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
385 # elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)):
386 # nn.init.constant_(m.weight, 1)
387 # nn.init.constant_(m.bias, 0)
388
389 def forward(self, x):
390 conv0 = self.conv0(x)
391 conv2 = self.conv2(self.conv1(conv0))
392 conv4 = self.conv4(self.conv3(conv2))
393 x = self.conv6(self.conv5(conv4))
394 x = conv4 + self.conv7(x)
395 x = conv2 + self.conv9(x)
396 x = conv0 + self.conv11(x)
397 x = self.prob(x)
398 return x
399
400class CostRegNet_part_refine(nn.Module):
401 def __init__(self, in_channels, base_channels,stage=0):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected