MCPcopy Create free account
hub / github.com/FreeformRobotics/OTS / build_encoder

Method build_encoder

models/models.py:62–78  ·  view source on GitHub ↗
(arch='resnet50dilated', fc_dim=512, weights='')

Source from the content-addressed store, hash-verified

60
61 @staticmethod
62 def build_encoder(arch='resnet50dilated', fc_dim=512, weights=''):
63 pretrained = True if len(weights) == 0 else False
64 arch = arch.lower()
65 if arch == 'resnet18dilated':
66 orig_resnet = resnet.__dict__['resnet18'](pretrained=pretrained)
67 net_encoder = ResnetDilated(orig_resnet, dilate_scale=8)
68 elif arch == 'resnet50dilated':
69 orig_resnet = resnet.__dict__['resnet50'](pretrained=pretrained)
70 net_encoder = ResnetDilated(orig_resnet, dilate_scale=8)
71 else:
72 raise Exception('Architecture undefined!')
73
74 if len(weights) > 0:
75 print('Loading weights for net_encoder')
76 net_encoder.load_state_dict(
77 torch.load(weights, map_location=lambda storage, loc: storage), strict=False)
78 return net_encoder
79
80 @staticmethod
81 def build_decoder(arch='ppm',

Callers 1

load_seg_moduleFunction · 0.80

Calls 1

ResnetDilatedClass · 0.85

Tested by

no test coverage detected