Linhz/ViMNer
1
1import torch.nn as nn
2import math
3import torch.utils.model_zoo as model_zoo
4
5
6__all__ = ['ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101',
7 'resnet152']
8
9
10model_urls = {
11 'resnet18': 'https://download.pytorch.org/models/resnet18-5c106cde.pth',
12 'resnet34': 'https://download.pytorch.org/models/resnet34-333f7ec4.pth',
13 'resnet50': 'https://download.pytorch.org/models/resnet50-19c8e357.pth',
14 'resnet101': 'https://download.pytorch.org/models/resnet101-5d3b4d8f.pth',
15 'resnet152': 'https://download.pytorch.org/models/resnet152-b121ed2d.pth',
16}
17
18
19def conv3x3(in_planes, out_planes, stride=1):
20 "3x3 convolution with padding"
21 return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,
22 padding=1, bias=False)
23
24
25class BasicBlock(nn.Module):
26 expansion = 1
27
28 def __init__(self, inplanes, planes, stride=1, downsample=None):
29 super(BasicBlock, self).__init__()
30 self.conv1 = conv3x3(inplanes, planes, stride)
31 self.bn1 = nn.BatchNorm2d(planes)
32 self.relu = nn.ReLU(inplace=True)
33 self.conv2 = conv3x3(planes, planes)
34 self.bn2 = nn.BatchNorm2d(planes)
35 self.downsample = downsample
36 self.stride = stride
37
38 def forward(self, x):
39 residual = x
40
41 out = self.conv1(x)
42 out = self.bn1(out)
43 out = self.relu(out)
44
45 out = self.conv2(out)
46 out = self.bn2(out)
47
48 if self.downsample is not None:
49 residual = self.downsample(x)
50
51 out += residual
52 out = self.relu(out)
53
54 return out
55
56
57class Bottleneck(nn.Module):
58 expansion = 4
59
60 def __init__(self, inplanes, planes, stride=1, downsample=None):
61 super(Bottleneck, self).__init__()
62 self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)
63 self.bn1 = nn.BatchNorm2d(planes)
64 self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride,
65 padding=1, bias=False)
66 self.bn2 = nn.BatchNorm2d(planes)
67 self.conv3 = nn.Conv2d(planes, planes * 4, kernel_size=1, bias=False)
68 self.bn3 = nn.BatchNorm2d(planes * 4)
69 self.relu = nn.ReLU(inplace=True)
70 self.downsample = downsample
71 self.stride = stride
72
73 def forward(self, x):
74 residual = x
75
76 out = self.conv1(x)
77 out = self.bn1(out)
78 out = self.relu(out)
79
80 out = self.conv2(out)
81 out = self.bn2(out)
82 out = self.relu(out)
83
84 out = self.conv3(out)
85 out = self.bn3(out)
86
87 if self.downsample is not None:
88 residual = self.downsample(x)
89
90 out += residual
91 out = self.relu(out)
92
93 return out
94
95
96class ResNet(nn.Module):
97
98 def __init__(self, block, layers, num_classes=1000):
99 self.inplanes = 64
100 super(ResNet, self).__init__()
101 self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3,
102 bias=False)
103 self.bn1 = nn.BatchNorm2d(64)
104 self.relu = nn.ReLU(inplace=True)
105 self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
106 self.layer1 = self._make_layer(block, 64, layers[0])
107 self.layer2 = self._make_layer(block, 128, layers[1], stride=2)
108 self.layer3 = self._make_layer(block, 256, layers[2], stride=2)
109 self.layer4 = self._make_layer(block, 512, layers[3], stride=2)
110 self.avgpool = nn.AvgPool2d(7, stride=1)
111 self.fc = nn.Linear(512 * block.expansion, num_classes)
112
113 for m in self.modules():
114 if isinstance(m, nn.Conv2d):
115 n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
116 m.weight.data.normal_(0, math.sqrt(2. / n))
117 elif isinstance(m, nn.BatchNorm2d):
118 m.weight.data.fill_(1)
119 m.bias.data.zero_()
120
121 def _make_layer(self, block, planes, blocks, stride=1):
122 downsample = None
123 if stride != 1 or self.inplanes != planes * block.expansion:
124 downsample = nn.Sequential(
125 nn.Conv2d(self.inplanes, planes * block.expansion,
126 kernel_size=1, stride=stride, bias=False),
127 nn.BatchNorm2d(planes * block.expansion),
128 )
129
130 layers = []
131 layers.append(block(self.inplanes, planes, stride, downsample))
132 self.inplanes = planes * block.expansion
133 for i in range(1, blocks):
134 layers.append(block(self.inplanes, planes))
135
136 return nn.Sequential(*layers)
137
138 def forward(self, x):
139 x = self.conv1(x)
140 x = self.bn1(x)
141 x = self.relu(x)
142 x = self.maxpool(x)
143
144 x = self.layer1(x)
145 x = self.layer2(x)
146 x = self.layer3(x)
147 x = self.layer4(x)
148
149 x = self.avgpool(x)
150 x = x.view(x.size(0), -1)
151 x = self.fc(x)
152
153 return x
154
155
156def resnet18(pretrained=False, **kwargs):
157 """Constructs a ResNet-18 model.
158
159 Args:
160 pretrained (bool): If True, returns a model pre-trained on ImageNet
161 """
162 model = ResNet(BasicBlock, [2, 2, 2, 2], **kwargs)
163 if pretrained:
164 model.load_state_dict(model_zoo.load_url(model_urls['resnet18']))
165 return model
166
167
168def resnet34(pretrained=False, **kwargs):
169 """Constructs a ResNet-34 model.
170
171 Args:
172 pretrained (bool): If True, returns a model pre-trained on ImageNet
173 """
174 model = ResNet(BasicBlock, [3, 4, 6, 3], **kwargs)
175 if pretrained:
176 model.load_state_dict(model_zoo.load_url(model_urls['resnet34']))
177 return model
178
179
180def resnet50(pretrained=False, **kwargs):
181 """Constructs a ResNet-50 model.
182
183 Args:
184 pretrained (bool): If True, returns a model pre-trained on ImageNet
185 """
186 model = ResNet(Bottleneck, [3, 4, 6, 3], **kwargs)
187 if pretrained:
188 model.load_state_dict(model_zoo.load_url(model_urls['resnet50']))
189 return model
190
191
192def resnet101(pretrained=False, **kwargs):
193 """Constructs a ResNet-101 model.
194
195 Args:
196 pretrained (bool): If True, returns a model pre-trained on ImageNet
197 """
198 model = ResNet(Bottleneck, [3, 4, 23, 3], **kwargs)
199 if pretrained:
200 model.load_state_dict(model_zoo.load_url(model_urls['resnet101']))
201 return model
202
203
204def resnet152(pretrained=False, **kwargs):
205 """Constructs a ResNet-152 model.
206
207 Args:
208 pretrained (bool): If True, returns a model pre-trained on ImageNet
209 """
210 model = ResNet(Bottleneck, [3, 8, 36, 3], **kwargs)
211 if pretrained:
212 model.load_state_dict(model_zoo.load_url(model_urls['resnet152']))
213 return model