MCPcopy Create free account
hub / github.com/cvg/NoPoSplat / evaluate

Function evaluate

src/eval_pose.py:45–72  ·  view source on GitHub ↗
(cfg_dict: DictConfig)

Source from the content-addressed store, hash-verified

43 config_name="main",
44)
45def evaluate(cfg_dict: DictConfig):
46 cfg = load_typed_config(cfg_dict, RootCfg,
47 {list[LossCfgWrapper]: separate_loss_cfg_wrappers,
48 list[DatasetCfgWrapper]: separate_dataset_cfg_wrappers},)
49 set_cfg(cfg_dict)
50 torch.manual_seed(cfg.seed)
51
52 encoder, encoder_visualizer = get_encoder(cfg.model.encoder)
53 ckpt_weights = torch.load(cfg.checkpointing.load, map_location='cpu')['state_dict']
54 # remove the prefix "encoder.", need to judge if is at start of key
55 ckpt_weights = {k[8:] if k.startswith("encoder.") else k: v for k, v in ckpt_weights.items()}
56 missing_keys, unexpected_keys = encoder.load_state_dict(ckpt_weights, strict=True)
57
58 trainer = Trainer(max_epochs=-1, accelerator="gpu", inference_mode=False)
59 pose_evaluator = PoseEvaluator(cfg.evaluation,
60 encoder,
61 get_decoder(cfg.model.decoder),
62 get_losses(cfg.loss))
63 data_module = DataModule(
64 cfg.dataset,
65 cfg.data_loader,
66 )
67
68 metrics = trainer.test(pose_evaluator, datamodule=data_module)
69
70 cfg.evaluation.output_metrics_path.parent.mkdir(exist_ok=True, parents=True)
71 with cfg.evaluation.output_metrics_path.open("w") as f:
72 json.dump(metrics[0], f)
73
74
75if __name__ == "__main__":

Callers 1

eval_pose.pyFile · 0.70

Calls 9

load_typed_configFunction · 0.90
set_cfgFunction · 0.90
get_encoderFunction · 0.90
PoseEvaluatorClass · 0.90
get_decoderFunction · 0.90
get_lossesFunction · 0.90
DataModuleClass · 0.90
dumpMethod · 0.80
load_state_dictMethod · 0.45

Tested by

no test coverage detected