| 400 | self.batch_generator = batch_generator |
| 401 | |
| 402 | def _make_model(self): |
| 403 | self.logger.info('Load checkpoint from {}'.format(cfg.pretrained_model_path)) |
| 404 | |
| 405 | # prepare network |
| 406 | self.logger.info("Creating graph...") |
| 407 | model = get_model('test') |
| 408 | model = DataParallel(model).cuda() |
| 409 | if not getattr(cfg, 'random_init', False): |
| 410 | ckpt = torch.load(cfg.pretrained_model_path, map_location=torch.device('cpu')) |
| 411 | |
| 412 | from collections import OrderedDict |
| 413 | new_state_dict = OrderedDict() |
| 414 | for k, v in ckpt['network'].items(): |
| 415 | if 'module' not in k: |
| 416 | k = 'module.' + k |
| 417 | k = k.replace('backbone', 'encoder').replace('body_rotation_net', 'body_regressor').replace( |
| 418 | 'hand_rotation_net', 'hand_regressor') |
| 419 | new_state_dict[k] = v |
| 420 | self.logger.warning("Attention: Strict=False is set for checkpoint loading. Please check manually.") |
| 421 | model.load_state_dict(new_state_dict, strict=False) |
| 422 | model.eval() |
| 423 | else: |
| 424 | print('Random init!!!!!!!') |
| 425 | |
| 426 | self.model = model |
| 427 | |
| 428 | def _evaluate(self, outs, cur_sample_idx): |
| 429 | eval_result = self.testset.evaluate(outs, cur_sample_idx) |