eval
dreem.inference.eval
¶
API for evaluating model against ground truth.
Functions:
| Name | Description |
|---|---|
run |
Run evaluation based on config. |
run(cfg)
¶
Run evaluation based on config.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
cfg
|
DictConfig
|
A DictConfig containing checkpoint path and data configuration |
required |
Source code in dreem/inference/eval.py
def run(cfg: DictConfig) -> None:
"""Run evaluation based on config.
Args:
cfg: A DictConfig containing checkpoint path and data configuration
"""
eval_cfg = Config(cfg)
checkpoint = eval_cfg.get("ckpt_path", None)
if checkpoint is None:
raise ValueError("Checkpoint path not found in config")
model = GTRRunner.load_from_checkpoint(checkpoint, strict=False)
overrides_dict = model.setup_tracking(eval_cfg, mode="eval")
logger.info(
f"Saving tracking results and metrics to {model.test_results['save_path']}"
)
labels_files, vid_files = eval_cfg.get_data_paths(
"test", eval_cfg.cfg.dataset.test_dataset
)
trainer = eval_cfg.get_trainer()
for label_file, vid_file in zip(labels_files, vid_files):
dataset = eval_cfg.get_dataset(
label_files=[label_file],
vid_files=[vid_file],
mode="test",
overrides=overrides_dict,
)
dataloader = eval_cfg.get_dataloader(dataset, mode="test")
_ = trainer.test(model, dataloader)
logger.info("Evaluation complete.")