MCPcopy Create free account
hub / github.com/DataScienceHamburg/PyTorchUltimateMaterial / SupervisedNet

Class SupervisedNet

350_SemiSupervised/super_learn.py:32–62  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

30
31#%% Model Class
32class SupervisedNet(nn.Module):
33 def __init__(self, n_super_classes) -> None:
34 super().__init__()
35 self.conv1 = nn.Conv2d(1, 6, 3)
36 self.pool = nn.MaxPool2d(2, 2)
37 self.conv2 = nn.Conv2d(6, 16, 3)
38 self.fc1 = nn.Linear(16 * 6 * 6, 128)
39 self.fc2 = nn.Linear(128, 64)
40 self.fc_out_super = nn.Linear(64, n_super_classes)
41 self.relu = nn.ReLU()
42 self.output_layer_super = nn.Sigmoid()
43
44 def backbone(self, x):
45 x = self.conv1(x)
46 x = self.relu(x)
47 x = self.pool(x)
48 x = self.conv2(x)
49 x = self.relu(x)
50 x = self.pool(x)
51 x = torch.flatten(x, 1)
52 x = self.fc1(x)
53 x = self.relu(x)
54 x = self.fc2(x)
55 x = self.relu(x)
56 return x
57
58 def forward(self, x):
59 x = self.backbone(x)
60 x = self.fc_out_super(x)
61 x = self.output_layer_super(x)
62 return x
63
64model = SupervisedNet(n_super_classes=2)
65model.train()

Callers 1

super_learn.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected