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