| 512 | |
| 513 | |
| 514 | class DilatedDCNNV2(nn.Module): |
| 515 | def __init__(self, in_channels, |
| 516 | filters, |
| 517 | multiplier=2, |
| 518 | n_dilation=4, |
| 519 | batch_norm=True, |
| 520 | activation=F.relu, |
| 521 | octave_conv=False): |
| 522 | super(DilatedDCNNV2, self).__init__() |
| 523 | assert not (octave_conv and batch_norm) |
| 524 | |
| 525 | self.n_dilation = n_dilation |
| 526 | self.conv_blocks = nn.ModuleList() |
| 527 | |
| 528 | |
| 529 | n_filter_per_module = (filters * 4) // n_dilation # multiplier |
| 530 | for dilation in range(n_dilation-1): |
| 531 | self.conv_blocks.append( |
| 532 | Conv3DConfigurable( |
| 533 | in_channels, |
| 534 | n_filter_per_module, |
| 535 | mid_filter=multiplier*filters, |
| 536 | dilation_rate=2 ** dilation, |
| 537 | use_bias=not batch_norm, |
| 538 | octave=octave_conv |
| 539 | ) |
| 540 | ) |
| 541 | self.conv_blocks.append( |
| 542 | Conv3DConfigurable( |
| 543 | in_channels, |
| 544 | (filters * 4) - n_filter_per_module * (n_dilation-1), # multiplier |
| 545 | mid_filter=multiplier*filters, |
| 546 | dilation_rate=2 ** (n_dilation-1), |
| 547 | use_bias=not batch_norm, |
| 548 | octave=octave_conv |
| 549 | ) |
| 550 | ) |
| 551 | |
| 552 | self.batch_norm = torch.nn.BatchNorm3d(num_features=filters * 4, eps=1e-3, momentum=0.1) if batch_norm else None |
| 553 | self.activation = activation |
| 554 | self.octave = octave_conv |
| 555 | |
| 556 | def forward(self, inputs): |
| 557 | x = [] |
| 558 | for block in self.conv_blocks: |
| 559 | x.append(block(inputs)) |
| 560 | x = torch.cat(x, dim=1) |
| 561 | |
| 562 | if self.octave: |
| 563 | raise Exception("Position octave 1: should not be here !!!") |
| 564 | |
| 565 | |
| 566 | if self.batch_norm is not None: |
| 567 | x = self.batch_norm(x) |
| 568 | |
| 569 | if self.activation is not None: |
| 570 | if self.octave: |
| 571 | raise Exception("Position octave 2: should not be here !!!") |