MCPcopy Create free account
hub / github.com/UVA-Computer-Vision-Lab/FrameINO / DilatedDCNNV2

Class DilatedDCNNV2

preprocess/auxiliary/AutoShot.py:514–574  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

512
513
514class 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 !!!")

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected