MCPcopy Create free account
hub / github.com/MotrixLab/ADHMR / _make_model

Method _make_model

HMR-Scorer/common/base.py:402–426  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 7

get_modelFunction · 0.85
printFunction · 0.85
infoMethod · 0.80
warningMethod · 0.80
loadMethod · 0.45
itemsMethod · 0.45
load_state_dictMethod · 0.45

Tested by

no test coverage detected