| 411 | |
| 412 | |
| 413 | class DilatedDCNNV2ABC(nn.Module): |
| 414 | def __init__(self, in_channels, |
| 415 | filters, |
| 416 | batch_norm=True, |
| 417 | activation=F.relu, |
| 418 | octave_conv=False, |
| 419 | multiplier=4, |
| 420 | n_dilation=4, |
| 421 | st_type="A"): |
| 422 | super(DilatedDCNNV2ABC, self).__init__() |
| 423 | assert not (octave_conv and batch_norm) |
| 424 | |
| 425 | self.share = torch.nn.Conv3d(in_channels=in_channels, |
| 426 | out_channels=multiplier * filters, |
| 427 | kernel_size=(1, 3, 3), |
| 428 | padding=(0, 1, 1), |
| 429 | dilation=(1, 1, 1), |
| 430 | bias=False) |
| 431 | init.kaiming_normal_(self.share.weight, mode="fan_in", nonlinearity="relu") |
| 432 | |
| 433 | self.conv_blocks = nn.ModuleList() |
| 434 | |
| 435 | n_in_plane = multiplier * filters |
| 436 | if st_type == "B": |
| 437 | n_in_plane = in_channels |
| 438 | |
| 439 | n_filter_per_module = (filters * 4) // n_dilation # multiplier |
| 440 | for dilation in range(n_dilation-1): |
| 441 | self.conv_blocks.append( |
| 442 | Conv3DConfigurable( |
| 443 | n_in_plane, |
| 444 | n_filter_per_module, |
| 445 | 2 ** dilation, |
| 446 | mid_filter=n_in_plane, |
| 447 | separable=True, |
| 448 | sharable=True, |
| 449 | use_bias=not batch_norm, |
| 450 | octave=octave_conv |
| 451 | ) |
| 452 | ) |
| 453 | self.conv_blocks.append( |
| 454 | Conv3DConfigurable( |
| 455 | n_in_plane, |
| 456 | (filters * 4) - n_filter_per_module * (n_dilation-1), # multiplier |
| 457 | 2 ** (n_dilation - 1), |
| 458 | mid_filter=n_in_plane, |
| 459 | separable=True, |
| 460 | sharable=True, |
| 461 | use_bias=not batch_norm, |
| 462 | octave=octave_conv |
| 463 | ) |
| 464 | ) |
| 465 | |
| 466 | self.octave = octave_conv |
| 467 | self.multiplier = multiplier |
| 468 | self.n_dilation = n_dilation |
| 469 | self.st_type = st_type |
| 470 | |