Evaluates given checkpoint on a datamodule testset.
This method is wrapped in optional @task_wrapper decorator, that controls the
behavior during failure. Useful for multiruns, saving info about the crash, etc.
Parameters:
| Name |
Type |
Description |
Default |
cfg |
DictConfig
|
DictConfig configuration composed by Hydra.
|
required
|
Returns:
Tuple[dict, dict] with metrics and dict with all instantiated objects.
Source code in meds_torch/eval.py
| @task_wrapper
def evaluate(cfg: DictConfig, datamodule=None) -> tuple[dict[str, Any], dict[str, Any]]:
"""Evaluates given checkpoint on a datamodule testset.
This method is wrapped in optional @task_wrapper decorator, that controls the
behavior during failure. Useful for multiruns, saving info about the crash, etc.
Args:
cfg: DictConfig configuration composed by Hydra.
Returns:
Tuple[dict, dict] with metrics and dict with all instantiated objects.
"""
assert cfg.ckpt_path
log.info(f"Instantiating datamodule <{cfg.data._target_}>")
if not datamodule:
datamodule: LightningDataModule = hydra.utils.instantiate(cfg.data)
log.info(f"Instantiating model <{cfg.model._target_}>")
model: LightningModule = hydra.utils.instantiate(cfg.model)
checkpoint = torch.load(cfg.ckpt_path, map_location="cpu", weights_only=False)
model.load_state_dict(checkpoint["state_dict"])
log.info("Instantiating loggers...")
logger: list[Logger] = instantiate_loggers(cfg.get("logger"))
log.info(f"Instantiating trainer <{cfg.trainer._target_}>")
trainer: Trainer = hydra.utils.instantiate(cfg.trainer, logger=logger)
object_dict = {
"cfg": cfg,
"datamodule": datamodule,
"model": model,
"logger": logger,
"trainer": trainer,
}
if logger:
log.info("Logging hyperparameters!")
log_hyperparameters(object_dict)
log.info("Starting testing!")
trainer.test(model=model, datamodule=datamodule)
# for predictions use trainer.predict(...)
# predictions = trainer.predict(model=model, dataloaders=dataloaders, ckpt_path=cfg.ckpt_path)
metric_dict = trainer.callback_metrics
return metric_dict, object_dict
|