ACL-LKNet / evaluate_ensemble.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw History Blame Contribute Delete
7.32 kB
#!/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()