stereoid's picture
Add files using upload-large-folder tool
af46737 verified
Raw
History Blame Contribute Delete
1.22 kB
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