| 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 | |
| 480 | def get_seg_model(): |