Image Classification
timm
English
medical-imaging
knee-mri
acl-tear-detection
deep-learning
convnext
self-attention
masked-slice-modeling
radiology
orthopedics
Eval Results (legacy)
Instructions to use shareefch1413/ACL-LKNet with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use shareefch1413/ACL-LKNet with timm:
import timm model = timm.create_model("hf-hub:shareefch1413/ACL-LKNet", pretrained=True) - Notebooks
- Google Colab
- Kaggle
Download evaluate_ensemble.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 7.32 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/evaluate_ensemble.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/evaluate_ensemble.py
-
curl -L -o evaluate_ensemble.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/evaluate_ensemble.py
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 | |
| 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() | |