MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / Bottleneck

Class Bottleneck

Image_Classification/src/models/mobilenetv3.py:84–157  ·  view source on GitHub ↗

The basic unit of MobileNetV3

Source from the content-addressed store, hash-verified

82
83
84class Bottleneck(nn.Module):
85 '''
86 The basic unit of MobileNetV3
87 '''
88
89 def __init__(self, in_channels_num, exp_size, out_channels_num, kernel_size, stride, use_SE, NL, BN_momentum):
90 '''
91 use_SE: True or False -- use SE Module or not
92 NL: nonlinearity, 'RE' or 'HS'
93 '''
94 super(Bottleneck, self).__init__()
95
96 assert stride in [1, 2]
97 NL = NL.upper()
98 assert NL in ['RE', 'HS']
99
100 use_HS = NL == 'HS'
101
102 # Whether to use residual structure or not
103 self.use_residual = (stride == 1 and in_channels_num == out_channels_num)
104
105 if exp_size == in_channels_num:
106 # Without expansion, the first depthwise convolution is omitted
107 self.conv1 = nn.Sequential(
108 # Depthwise Convolution
109 nn.Conv2d(in_channels=in_channels_num, out_channels=exp_size, kernel_size=kernel_size, stride=stride,
110 padding=(kernel_size - 1) // 2, bias=False),
111 nn.BatchNorm2d(num_features=exp_size, momentum=BN_momentum),
112 # SE Module
113 SEModule(exp_size) if use_SE else nn.Sequential(),
114 H_swish() if use_HS else nn.ReLU(inplace=False))
115 self.conv2 = nn.Sequential(
116 # Linear Pointwise Convolution
117 nn.Conv2d(in_channels=exp_size, out_channels=out_channels_num, kernel_size=1, stride=1, padding=0,
118 bias=False),
119 # nn.BatchNorm2d(num_features=out_channels_num, momentum=BN_momentum)
120 nn.Sequential(
121 OrderedDict([('lastBN', nn.BatchNorm2d(num_features=out_channels_num))])) if self.use_residual else
122 nn.BatchNorm2d(num_features=out_channels_num, momentum=BN_momentum)
123 )
124 else:
125 # With expansion
126 self.conv1 = nn.Sequential(
127 # Pointwise Convolution for expansion
128 nn.Conv2d(in_channels=in_channels_num, out_channels=exp_size, kernel_size=1, stride=1, padding=0,
129 bias=False),
130 nn.BatchNorm2d(num_features=exp_size, momentum=BN_momentum),
131 H_swish() if use_HS else nn.ReLU(inplace=False))
132 self.conv2 = nn.Sequential(
133 # Depthwise Convolution
134 nn.Conv2d(in_channels=exp_size, out_channels=exp_size, kernel_size=kernel_size, stride=stride,
135 padding=(kernel_size - 1) // 2, bias=False),
136 nn.BatchNorm2d(num_features=exp_size, momentum=BN_momentum),
137 # SE Module
138 SEModule(exp_size) if use_SE else nn.Sequential(),
139 H_swish() if use_HS else nn.ReLU(inplace=False),
140 # Linear Pointwise Convolution
141 nn.Conv2d(in_channels=exp_size, out_channels=out_channels_num, kernel_size=1, stride=1, padding=0,

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected