(cfg_dict: DictConfig)
| 43 | config_name="main", |
| 44 | ) |
| 45 | def 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 | |
| 75 | if __name__ == "__main__": |
no test coverage detected