| 30 | |
| 31 | #%% Model Class |
| 32 | class 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 | |
| 64 | model = SupervisedNet(n_super_classes=2) |
| 65 | model.train() |