| import torch.nn as nn | |
| import torch.nn.functional as F | |
| import torchvision.models as models | |
| NUM_CLASSES = { | |
| 'interaction': 54, | |
| 'semantics': 771 | |
| } | |
| class CNN(nn.Module): | |
| def __init__(self, task, softmax = True): | |
| super(CNN, self).__init__() | |
| assert task in NUM_CLASSES.keys(), "Task must be one of {}".format(NUM_CLASSES.keys()) | |
| self.task = task | |
| self.softmax = softmax | |
| self.num_classes = NUM_CLASSES[task] | |
| self.resnet = models.resnet50(weights = models.ResNet50_Weights.DEFAULT) | |
| num_features = self.resnet.fc.in_features | |
| self.resnet.fc = nn.Sequential( | |
| nn.Linear(num_features, 128), | |
| nn.ReLU(), | |
| nn.Dropout(0.5), | |
| nn.Linear(128, self.num_classes) | |
| ) | |
| self._freeze_conv() | |
| def forward(self, x): | |
| x = self.resnet(x) | |
| if self.softmax: | |
| x = F.softmax(x, dim = 1) | |
| return x | |
| def _freeze_conv(self): | |
| for param in self.resnet.parameters(): | |
| param.requires_grad = False | |
| for param in self.resnet.fc.parameters(): | |
| param.requires_grad = True | |