| 59 | |
| 60 | |
| 61 | class MobileNetV2(nn.Module): |
| 62 | def __init__(self, num_classes=1000, alpha=1.0, round_nearest=8): |
| 63 | super(MobileNetV2, self).__init__() |
| 64 | block = InvertedResidual |
| 65 | input_channel = _make_divisible(32 * alpha, round_nearest) |
| 66 | last_channel = _make_divisible(1280 * alpha, round_nearest) |
| 67 | |
| 68 | inverted_residual_setting = [ |
| 69 | # t,c,n,s |
| 70 | [1, 16, 1, 1], |
| 71 | [6, 24, 2, 2], |
| 72 | [6, 32, 3, 2], |
| 73 | [6, 64, 4, 2], |
| 74 | [6, 96, 3, 1], |
| 75 | [6, 160, 3, 2], |
| 76 | [6, 320, 1, 1], |
| 77 | ] |
| 78 | |
| 79 | features = [] |
| 80 | # conv1 layer |
| 81 | features.append(ConvBNReLU(3, input_channel, stride=2)) |
| 82 | # building inverted residual blockes |
| 83 | for t, c, n, s in inverted_residual_setting: |
| 84 | output_channel = _make_divisible(c * alpha, round_nearest) |
| 85 | for i in range(n): |
| 86 | stride = s if i == 0 else 1 |
| 87 | features.append(block(input_channel, output_channel, stride, expand_ratio=t)) |
| 88 | input_channel = output_channel |
| 89 | # building last several layers |
| 90 | features.append(ConvBNReLU(input_channel, last_channel, 1)) |
| 91 | # combine features layers |
| 92 | self.features = nn.Sequential(*features) |
| 93 | |
| 94 | # building classifier |
| 95 | self.avgpool = nn.AdaptiveAvgPool2d((1, 1)) |
| 96 | self.classifier = nn.Sequential( |
| 97 | nn.Dropout(0.2), |
| 98 | nn.Linear(last_channel, num_classes) |
| 99 | ) |
| 100 | |
| 101 | # weight initialization |
| 102 | for m in self.modules(): |
| 103 | if isinstance(m, nn.Conv2d): |
| 104 | nn.init.kaiming_normal_(m.weight, mode='fan_out') |
| 105 | if m.bias is not None: |
| 106 | nn.init.zeros_(m.bias) |
| 107 | elif isinstance(m, nn.BatchNorm2d): |
| 108 | nn.init.ones_(m.weight) |
| 109 | nn.init.zeros_(m.bias) |
| 110 | elif isinstance(m, nn.Linear): |
| 111 | nn.init.normal_(m.weight, 0, 0.01) |
| 112 | nn.init.zeros_(m.bias) |
| 113 | |
| 114 | |
| 115 | def forward(self, x): |
| 116 | x = self.features(x) |
| 117 | x = self.avgpool(x) |
| 118 | x = torch.flatten(x, 1) |