MCPcopy Create free account
hub / github.com/NVIDIA/semantic-segmentation / init_weights

Method init_weights

network/hrnetv2.py:451–477  ·  view source on GitHub ↗
(self, pretrained=cfg.MODEL.HRNET_CHECKPOINT)

Source from the content-addressed store, hash-verified

449 return None, None, feats
450
451 def init_weights(self, pretrained=cfg.MODEL.HRNET_CHECKPOINT):
452 logx.msg('=> init weights from normal distribution')
453 for name, m in self.named_modules():
454 if any(part in name for part in {'cls', 'aux', 'ocr'}):
455 # print('skipped', name)
456 continue
457 if isinstance(m, nn.Conv2d):
458 nn.init.normal_(m.weight, std=0.001)
459 elif isinstance(m, cfg.MODEL.BNFUNC):
460 nn.init.constant_(m.weight, 1)
461 nn.init.constant_(m.bias, 0)
462 if os.path.isfile(pretrained):
463 pretrained_dict = torch.load(pretrained,
464 map_location={'cuda:0': 'cpu'})
465 logx.msg('=> loading pretrained model {}'.format(pretrained))
466 model_dict = self.state_dict()
467 pretrained_dict = {k.replace('last_layer',
468 'aux_head').replace('model.', ''): v
469 for k, v in pretrained_dict.items()}
470 #print(set(model_dict) - set(pretrained_dict))
471 #print(set(pretrained_dict) - set(model_dict))
472 pretrained_dict = {k: v for k, v in pretrained_dict.items()
473 if k in model_dict.keys()}
474 model_dict.update(pretrained_dict)
475 self.load_state_dict(model_dict)
476 elif pretrained:
477 raise RuntimeError('No such file {}'.format(pretrained))
478
479
480def get_seg_model():

Callers 1

get_seg_modelFunction · 0.95

Calls 1

updateMethod · 0.80

Tested by

no test coverage detected