MCPcopy Create free account
hub / github.com/LAMDA-CL/CVPR22-Fact / __init__

Method __init__

models/fact/Network.py:12–39  ·  view source on GitHub ↗
(self, args, mode=None)

Source from the content-addressed store, hash-verified

10class MYNET(nn.Module):
11
12 def __init__(self, args, mode=None):
13 super().__init__()
14
15 self.mode = mode
16 self.args = args
17 if self.args.dataset in ['cifar100','manyshotcifar']:
18 self.encoder = resnet20()
19 self.num_features = 64
20 if self.args.dataset in ['mini_imagenet','manyshotmini','imagenet100','imagenet1000']:
21 self.encoder = resnet18(False, args) # pretrained=False
22 self.num_features = 512
23 if self.args.dataset == 'cub200':
24 self.encoder = resnet18(True, args) # pretrained=True follow TOPIC, models for cub is imagenet pre-trained. https://github.com/xyutao/fscil/issues/11#issuecomment-687548790
25 self.num_features = 512
26 self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
27
28
29 self.pre_allocate = self.args.num_classes
30 self.fc = nn.Linear(self.num_features, self.pre_allocate, bias=False)
31
32 nn.init.orthogonal_(self.fc.weight)
33 self.dummy_orthogonal_classifier=nn.Linear(self.num_features, self.pre_allocate-self.args.base_class, bias=False)
34 self.dummy_orthogonal_classifier.weight.requires_grad = False
35
36 self.dummy_orthogonal_classifier.weight.data=self.fc.weight.data[self.args.base_class:,:]
37 print(self.dummy_orthogonal_classifier.weight.data.size())
38
39 print('self.dummy_orthogonal_classifier.weight initialized over.')
40
41 def forward_metric(self, x):
42 x = self.encode(x)

Callers

nothing calls this directly

Calls 2

resnet20Function · 0.85
resnet18Function · 0.85

Tested by

no test coverage detected