| 399 | |
| 400 | class 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 | |