The basic unit of MobileNetV3
| 82 | |
| 83 | |
| 84 | class 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, |