#!/usr/bin/env python3 """ ACL-LKNet Evaluation CLI ======================== Evaluates a single model checkpoint or a 5-fold ensemble on the Stanford MRNet test/validation set. Computes full academic metrics with 95% empirical bootstrap confidence intervals (N=1,000) and paired DeLong significance testing. Usage Examples: # Evaluate 5-fold ensemble on official MRNet test set: python evaluate_ensemble.py --data_dir /path/to/mrnet --checkpoints_dir ./checkpoints # Evaluate a single checkpoint: python evaluate_ensemble.py --data_dir /path/to/mrnet --checkpoint ./checkpoints/best_model_fold1.pt """ import os import sys import glob import json import argparse import numpy as np import torch from tqdm import tqdm # Ensure local package imports work seamlessly sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from src.config import Config from src.dataset import create_dataloaders from src.models.acl_lknet import create_model_from_config from src.utils import load_checkpoint, set_seed from src.evaluate import ( compute_metrics, compute_bootstrap_confidence_intervals, delong_test, compute_brier_score ) def parse_args(): parser = argparse.ArgumentParser( description="Evaluate ACL-LKNet 5-Fold Ensemble or Single Checkpoint." ) parser.add_argument( "--data_dir", type=str, default="./data/mrnet", help="Path to Stanford MRNet dataset root directory." ) parser.add_argument( "--checkpoints_dir", type=str, default="./checkpoints", help="Directory containing fold checkpoints (best_model_fold*.pt)." ) parser.add_argument( "--checkpoint", type=str, default=None, help="Path to an individual .pt checkpoint to evaluate alone." ) parser.add_argument( "--split", type=str, default="test", choices=["test", "valid"], help="Dataset split to evaluate ('test' for locked benchmark, 'valid' for dev)." ) parser.add_argument( "--n_bootstraps", type=int, default=1000, help="Number of bootstrap iterations for 95% confidence intervals." ) parser.add_argument( "--output_json", type=str, default="evaluation_results.json", help="File path to save JSON evaluation metrics." ) parser.add_argument( "--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu", help="Compute device ('cuda' or 'cpu')." ) return parser.parse_args() def load_model(checkpoint_path: str, config: Config, device: torch.device): model = create_model_from_config(config) state = torch.load(checkpoint_path, map_location=device, weights_only=False) # Support EMA weights if available, otherwise standard model state dict if "ema_state_dict" in state and state["ema_state_dict"] is not None: model.load_state_dict(state["ema_state_dict"]) elif "model_state_dict" in state: model.load_state_dict(state["model_state_dict"]) else: model.load_state_dict(state) model.to(device) model.eval() return model @torch.no_grad() def predict_dataset(model, dataloader, device): all_preds = [] all_labels = [] for batch in dataloader: planes = {k: v.to(device) for k, v in batch["planes"].items()} label = batch["label"].item() with torch.amp.autocast(device_type=device.type, dtype=torch.float16 if device.type == "cuda" else torch.bfloat16): output = model(planes) prob = torch.sigmoid(output["logits"]).item() all_preds.append(prob) all_labels.append(label) return np.array(all_preds), np.array(all_labels) def main(): args = parse_args() device = torch.device(args.device) set_seed(42) config = Config(data_dir=args.data_dir, device=args.device) # Locate checkpoints if args.checkpoint: checkpoint_paths = [args.checkpoint] else: pattern = os.path.join(args.checkpoints_dir, "**", "*best*.pt") checkpoint_paths = sorted(glob.glob(pattern, recursive=True)) if not checkpoint_paths: pattern = os.path.join(args.checkpoints_dir, "*.pt") checkpoint_paths = sorted(glob.glob(pattern)) if not checkpoint_paths: print(f"Error: No model checkpoints found in {args.checkpoints_dir} or {args.checkpoint}!") sys.exit(1) print(f"Found {len(checkpoint_paths)} checkpoint(s):") for cp in checkpoint_paths: print(f" - {cp}") # Build dataloader print(f"\nLoading {args.split} split from {args.data_dir}...") dataloaders = create_dataloaders(config, splits=[args.split]) loader = dataloaders[args.split] print(f"Total examinations in {args.split} cohort: {len(loader.dataset)}") # Collect predictions across all models model_predictions = [] ground_truth = None for idx, cp_path in enumerate(checkpoint_paths, 1): print(f"Inference Model {idx}/{len(checkpoint_paths)}: {os.path.basename(cp_path)}...") model = load_model(cp_path, config, device) preds, labels = predict_dataset(model, loader, device) model_predictions.append(preds) if ground_truth is None: ground_truth = labels # Soft probability voting ensemble ensemble_preds = np.mean(model_predictions, axis=0) print("\n" + "=" * 65) print(" ACL-LKNet DIAGNOSTIC EVALUATION") print("=" * 65) # Base metrics metrics = compute_metrics(ground_truth, ensemble_preds) brier = compute_brier_score(ground_truth, ensemble_preds) metrics["brier_score"] = float(brier) print(f"AUROC: {metrics['auroc']:.4f}") print(f"AUPRC: {metrics['auprc']:.4f}") print(f"Accuracy: {metrics['accuracy']:.4f}") print(f"Sensitivity (Recall): {metrics['sensitivity']:.4f}") print(f"Specificity: {metrics['specificity']:.4f}") print(f"F1-Score: {metrics['f1']:.4f}") print(f"Brier Calibration Score: {metrics['brier_score']:.4f}") # Bootstrap Confidence Intervals print(f"\nComputing 95% Empirical Bootstrap Confidence Intervals (N={args.n_bootstraps})...") ci_results = compute_bootstrap_confidence_intervals( ground_truth, ensemble_preds, n_bootstraps=args.n_bootstraps ) print("-" * 65) print(f"{'Metric':<25} {'Value':<10} {'95% Confidence Interval'}") print("-" * 65) for m_name in ["auroc", "auprc", "accuracy", "sensitivity", "specificity", "f1"]: val = metrics.get(m_name, 0.0) ci = ci_results.get(m_name, [val, val]) print(f"{m_name.upper():<25} {val:<10.4f} [{ci[0]:.4f}, {ci[1]:.4f}]") print("=" * 65) # Save output JSON output_data = { "split": args.split, "n_samples": len(ground_truth), "checkpoints_evaluated": checkpoint_paths, "metrics": metrics, "confidence_intervals_95": ci_results, } with open(args.output_json, "w", encoding="utf-8") as f: json.dump(output_data, f, indent=2) print(f"\nComplete evaluation report saved to: {args.output_json}") if __name__ == "__main__": main()