MCPcopy Create free account
hub / github.com/VisionLearningGroup/OVANet / ResBase

Class ResBase

models/basenet.py:7–39  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5
6
7class ResBase(nn.Module):
8 def __init__(self, option='resnet50', pret=True, top=False):
9 super(ResBase, self).__init__()
10 self.dim = 2048
11 self.top = top
12 if option == 'resnet18':
13 model_ft = models.resnet18(pretrained=pret)
14 self.dim = 512
15 if option == 'resnet34':
16 model_ft = models.resnet34(pretrained=pret)
17 self.dim = 512
18 if option == 'resnet50':
19 model_ft = models.resnet50(pretrained=pret)
20 if option == 'resnet101':
21 model_ft = models.resnet101(pretrained=pret)
22 if option == 'resnet152':
23 model_ft = models.resnet152(pretrained=pret)
24
25 if top:
26 self.features = model_ft
27 else:
28 mod = list(model_ft.children())
29 mod.pop()
30 self.features = nn.Sequential(*mod)
31
32
33 def forward(self, x):
34 x = self.features(x)
35 if self.top:
36 return x
37 else:
38 x = x.view(x.size(0), self.dim)
39 return x
40
41
42class VGGBase(nn.Module):

Callers 1

get_model_mmeFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected