| 356 | results=torch.cat((self.cosR_small(x),self.cosR_huge(x)),axis=1) |
| 357 | return results |
| 358 | class 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 | |
| 400 | class CostRegNet_part_refine(nn.Module): |
| 401 | def __init__(self, in_channels, base_channels,stage=0): |