AlexN/pull_up
1
1# -*- coding: utf-8 -*-2"""3Created on Sun Jul 4 15:07:27 20214 5@author: AlexandreN6"""7from __future__ import print_function, division8 9import torch10import torch.nn as nn11import torchvision12 13 14class SingleTractionHead(nn.Module):15 16 def __init__(self):17 super(SingleTractionHead, self).__init__()18 19 self.head_locs = nn.Sequential(nn.Linear(2048, 1024),20 nn.ReLU(), 21 nn.Dropout(p=0.3),22 nn.Linear(1024, 4),23 nn.Sigmoid()24 )25 26 # Head class should output the logits over the classe27 self.head_class = nn.Sequential(nn.Linear(2048, 128),28 nn.ReLU(), 29 nn.Dropout(p=0.3),30 nn.Linear(128, 1))31 32 def forward(self, features):33 features = features.view(features.size()[0], -1)34 35 y_bbox = self.head_locs(features)36 y_class = self.head_class(features)37 38 res = (y_bbox, y_class)39 return res40 41 42def create_model():43 # setup the architecture of the model44 feature_extractor = torchvision.models.resnet50(pretrained=True)45 model_body = nn.Sequential(*list(feature_extractor.children())[:-1])46 for param in model_body.parameters():47 param.requires_grad = False48 # Parameters of newly constructed modules have requires_grad=True by default49 # num_ftrs = model_body.fc.in_features50 51 model_head = SingleTractionHead()52 model = nn.Sequential(model_body, model_head)53 return model54 55 56def load_weights(model, path='model.pt', device_='cpu'):57 checkpoint = torch.load(path, map_location=torch.device(device_))58 model.load_state_dict(checkpoint)59 return model60 