File size: 7,317 Bytes
00801a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
#!/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()