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 src/evaluate.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 45.9 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/evaluate.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/evaluate.py
45.9 kB
| """ | |
| Evaluation module for ACL-LKNet. | |
| Includes: | |
| 1. Comprehensive Classification Performance: | |
| - AUROC, AUPRC, Accuracy, Balanced Accuracy, Sensitivity (Recall), Specificity, | |
| PPV (Precision), NPV, F1, MCC, Diagnostic Odds Ratio (DOR), Type I & II errors. | |
| - 95% Bootstrap Confidence Intervals (N=1000). | |
| 2. Thresholding-Based Evaluation: | |
| - Default (0.5), Youden's Index (J), F1-Optimal, High-Sensitivity Screening (Recall >= 95%). | |
| - Comprehensive threshold sweep tables and operating curves. | |
| 3. Confusion Matrix Analysis (Both): | |
| - Side-by-side Dual Plots: Raw Integer Counts + Condition-Normalized Percentages. | |
| 4. Quantitative Explainability Evaluation: | |
| - Anatomical Cruciate Landmark Pointing Game Hit Rate. | |
| - Attention Mass Concentration & Slice Attention Entropy / Sparsity. | |
| - Perturbation Faithfulness (Impact drop on masking top-attended slices). | |
| 5. Segmentation & Localization Architecture Evaluation: | |
| - Slice localization Pseudo-Dice and IoU against central ligament slices. | |
| - Masked Slice Modeling (MSM) pretext reconstruction fidelity (MSE, PSNR). | |
| - Architectural parameter and efficiency benchmark table. | |
| 6. Cross-Dataset Generalization & Robustness Testing: | |
| - Scanner perturbation stress tests (Rician noise, slice thickness, B1 field bias). | |
| - External dataset evaluation interface. | |
| 7. Paired Statistical Hypothesis Testing: | |
| - DeLong's test for AUROC difference significance. | |
| - McNemar's test for paired accuracy significance. | |
| """ | |
| import os | |
| import math | |
| import copy | |
| from typing import Dict, Optional, Tuple, List, Any | |
| import numpy as np | |
| import pandas as pd | |
| import matplotlib | |
| matplotlib.use("Agg") # Non-interactive backend for server/notebook execution | |
| import matplotlib.pyplot as plt | |
| import seaborn as sns | |
| from sklearn.metrics import ( | |
| roc_auc_score, | |
| average_precision_score, | |
| accuracy_score, | |
| balanced_accuracy_score, | |
| f1_score, | |
| confusion_matrix, | |
| roc_curve, | |
| precision_recall_curve, | |
| matthews_corrcoef, | |
| ) | |
| from scipy import stats | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| # ββ Core Diagnostic Contingency Helper ββββββββββββββββββββββββββββββ | |
| def _calc_contingency(labels: np.ndarray, predictions: np.ndarray, threshold: float = 0.5) -> Tuple[int, int, int, int]: | |
| """Return raw counts: (tn, fp, fn, tp).""" | |
| binary_preds = (predictions >= threshold).astype(int) | |
| cm = confusion_matrix(labels, binary_preds, labels=[0, 1]) | |
| if cm.shape == (2, 2): | |
| tn, fp, fn, tp = cm.ravel() | |
| else: | |
| tn = int(np.sum((labels == 0) & (binary_preds == 0))) | |
| fp = int(np.sum((labels == 0) & (binary_preds == 1))) | |
| fn = int(np.sum((labels == 1) & (binary_preds == 0))) | |
| tp = int(np.sum((labels == 1) & (binary_preds == 1))) | |
| return int(tn), int(fp), int(fn), int(tp) | |
| # ββ Comprehensive Classification Performance βββββββββββββββββββββββββ | |
| def compute_metrics( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| threshold: float = 0.5, | |
| prefix: str = "", | |
| ) -> Dict[str, float]: | |
| """ | |
| Compute comprehensive clinical diagnostic classification metrics. | |
| Returns: | |
| AUROC, AUPRC, Accuracy, Balanced_Accuracy, Sensitivity (Recall), Specificity, | |
| PPV (Precision), NPV, F1, MCC, DOR, Type1_Error_Rate, Type2_Error_Rate. | |
| """ | |
| labels = np.asarray(labels).astype(int) | |
| predictions = np.asarray(predictions).astype(float) | |
| binary_preds = (predictions >= threshold).astype(int) | |
| p = prefix + "_" if prefix and not prefix.endswith("_") else prefix | |
| tn, fp, fn, tp = _calc_contingency(labels, predictions, threshold) | |
| total = tn + fp + fn + tp | |
| positives = tp + fn | |
| negatives = tn + fp | |
| metrics: Dict[str, float] = {} | |
| # Rank metrics (threshold-independent) | |
| unique_labels = np.unique(labels) | |
| if len(unique_labels) < 2: | |
| metrics[f"{p}auroc"] = 0.0 | |
| metrics[f"{p}auprc"] = 0.0 | |
| else: | |
| try: | |
| metrics[f"{p}auroc"] = float(roc_auc_score(labels, predictions)) | |
| metrics[f"{p}auprc"] = float(average_precision_score(labels, predictions)) | |
| except Exception: | |
| metrics[f"{p}auroc"] = 0.0 | |
| metrics[f"{p}auprc"] = 0.0 | |
| # Accuracy & Balanced Accuracy | |
| metrics[f"{p}accuracy"] = float(accuracy_score(labels, binary_preds)) | |
| sens = tp / positives if positives > 0 else 0.0 | |
| spec = tn / negatives if negatives > 0 else 0.0 | |
| metrics[f"{p}balanced_accuracy"] = float(0.5 * (sens + spec)) | |
| # Diagnostic Sensitivity (Recall) & Specificity | |
| metrics[f"{p}sensitivity"] = float(sens) | |
| metrics[f"{p}specificity"] = float(spec) | |
| # Predictive Values (PPV & NPV) | |
| ppv = tp / (tp + fp) if (tp + fp) > 0 else 0.0 | |
| npv = tn / (tn + fn) if (tn + fn) > 0 else 0.0 | |
| metrics[f"{p}ppv"] = float(ppv) | |
| metrics[f"{p}npv"] = float(npv) | |
| # F1 Score | |
| metrics[f"{p}f1"] = float(f1_score(labels, binary_preds, zero_division=0)) | |
| # Matthews Correlation Coefficient (MCC) | |
| try: | |
| metrics[f"{p}mcc"] = float(matthews_corrcoef(labels, binary_preds)) | |
| except Exception: | |
| metrics[f"{p}mcc"] = 0.0 | |
| # Diagnostic Odds Ratio (DOR) with Haldane-Anscombe continuity correction (+0.5) | |
| dor = ((tp + 0.5) * (tn + 0.5)) / ((fp + 0.5) * (fn + 0.5)) | |
| metrics[f"{p}dor"] = float(dor) | |
| # Clinical Error Rates | |
| metrics[f"{p}type1_error_rate"] = float(fp / negatives) if negatives > 0 else 0.0 # False Positive Rate (alpha) | |
| metrics[f"{p}type2_error_rate"] = float(fn / positives) if positives > 0 else 0.0 # False Negative Rate (beta) | |
| return metrics | |
| # ββ Bootstrap Confidence Intervals ββββββββββββββββββββββββββββββββββ | |
| def bootstrap_ci( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| metric_fn, | |
| n_bootstrap: int = 1000, | |
| confidence: float = 0.95, | |
| seed: int = 42, | |
| ) -> Tuple[float, float, float]: | |
| """ | |
| Compute bootstrap confidence interval for a metric. | |
| Returns: | |
| (point_estimate, lower_bound, upper_bound) | |
| """ | |
| labels = np.asarray(labels) | |
| predictions = np.asarray(predictions) | |
| rng = np.random.RandomState(seed) | |
| n = len(labels) | |
| point = float(metric_fn(labels, predictions)) | |
| scores = [] | |
| for _ in range(n_bootstrap): | |
| idx = rng.choice(n, size=n, replace=True) | |
| try: | |
| score = metric_fn(labels[idx], predictions[idx]) | |
| if not math.isnan(score): | |
| scores.append(score) | |
| except (ValueError, ZeroDivisionError): | |
| continue | |
| if len(scores) < 10: | |
| return point, point, point | |
| alpha = 1 - confidence | |
| lower = float(np.percentile(scores, 100 * alpha / 2)) | |
| upper = float(np.percentile(scores, 100 * (1 - alpha / 2))) | |
| return point, lower, upper | |
| def compute_metrics_with_ci( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| threshold: float = 0.5, | |
| n_bootstrap: int = 1000, | |
| confidence: float = 0.95, | |
| ) -> Dict[str, Tuple[float, float, float]]: | |
| """ | |
| Compute all primary diagnostic metrics with empirical 95% bootstrap CIs. | |
| Returns: | |
| Dict mapping metric_name -> (point_estimate, lower_bound, upper_bound) | |
| """ | |
| labels = np.asarray(labels) | |
| predictions = np.asarray(predictions) | |
| metric_fns = { | |
| "AUROC": lambda y, p: roc_auc_score(y, p) if len(np.unique(y)) > 1 else 0.0, | |
| "AUPRC": lambda y, p: average_precision_score(y, p) if len(np.unique(y)) > 1 else 0.0, | |
| "Accuracy": lambda y, p: accuracy_score(y, (p >= threshold).astype(int)), | |
| "Balanced_Accuracy": lambda y, p: balanced_accuracy_score(y, (p >= threshold).astype(int)), | |
| "Sensitivity": lambda y, p: _sensitivity(y, p, threshold), | |
| "Specificity": lambda y, p: _specificity(y, p, threshold), | |
| "PPV": lambda y, p: _ppv(y, p, threshold), | |
| "NPV": lambda y, p: _npv(y, p, threshold), | |
| "F1": lambda y, p: f1_score(y, (p >= threshold).astype(int), zero_division=0), | |
| "MCC": lambda y, p: matthews_corrcoef(y, (p >= threshold).astype(int)), | |
| "DOR": lambda y, p: _dor(y, p, threshold), | |
| } | |
| results = {} | |
| for name, fn in metric_fns.items(): | |
| try: | |
| point, lower, upper = bootstrap_ci( | |
| labels, predictions, fn, n_bootstrap=n_bootstrap, confidence=confidence | |
| ) | |
| results[name] = (point, lower, upper) | |
| except Exception: | |
| results[name] = (0.0, 0.0, 0.0) | |
| return results | |
| def _sensitivity(labels, predictions, threshold=0.5): | |
| tn, fp, fn, tp = _calc_contingency(labels, predictions, threshold) | |
| return tp / (tp + fn) if (tp + fn) > 0 else 0.0 | |
| def _specificity(labels, predictions, threshold=0.5): | |
| tn, fp, fn, tp = _calc_contingency(labels, predictions, threshold) | |
| return tn / (tn + fp) if (tn + fp) > 0 else 0.0 | |
| def _ppv(labels, predictions, threshold=0.5): | |
| tn, fp, fn, tp = _calc_contingency(labels, predictions, threshold) | |
| return tp / (tp + fp) if (tp + fp) > 0 else 0.0 | |
| def _npv(labels, predictions, threshold=0.5): | |
| tn, fp, fn, tp = _calc_contingency(labels, predictions, threshold) | |
| return tn / (tn + fn) if (tn + fn) > 0 else 0.0 | |
| def _dor(labels, predictions, threshold=0.5): | |
| tn, fp, fn, tp = _calc_contingency(labels, predictions, threshold) | |
| return float(((tp + 0.5) * (tn + 0.5)) / ((fp + 0.5) * (fn + 0.5))) | |
| # ββ Thresholding-Based Evaluation Engine βββββββββββββββββββββββββββββ | |
| def find_optimal_thresholds( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| target_sensitivity: float = 0.95, | |
| ) -> Dict[str, Dict[str, float]]: | |
| """ | |
| Determine clinically relevant operating thresholds: | |
| 1. Default threshold (tau = 0.5) | |
| 2. Youden's J Index: max(Sensitivity + Specificity - 1) | |
| 3. F1-Optimal threshold: max F1 score | |
| 4. High-Sensitivity Screening: minimum threshold achieving Recall >= target_sensitivity (e.g. 0.95) | |
| Returns: | |
| Dict mapping mode -> {'threshold': tau, ...metrics} | |
| """ | |
| labels = np.asarray(labels).astype(int) | |
| predictions = np.asarray(predictions).astype(float) | |
| # Threshold candidates from ROC curve | |
| fpr, tpr, thresholds = roc_curve(labels, predictions) | |
| thresholds = np.clip(thresholds, 0.01, 0.99) | |
| # 1. Youden's Index J = TPR - FPR | |
| j_scores = tpr - fpr | |
| best_j_idx = int(np.argmax(j_scores)) | |
| youden_thresh = float(thresholds[best_j_idx]) | |
| # 2. F1-optimal & Screening thresholds via fine sweep | |
| sweep_taus = np.linspace(0.01, 0.99, 200) | |
| f1_list = [] | |
| sens_list = [] | |
| for tau in sweep_taus: | |
| b = (predictions >= tau).astype(int) | |
| f1_list.append(f1_score(labels, b, zero_division=0)) | |
| cm = confusion_matrix(labels, b, labels=[0, 1]) | |
| s = cm[1, 1] / (cm[1, 1] + cm[1, 0]) if (cm[1, 1] + cm[1, 0]) > 0 else 0.0 | |
| sens_list.append(s) | |
| best_f1_idx = int(np.argmax(f1_list)) | |
| f1_thresh = float(sweep_taus[best_f1_idx]) | |
| # High sensitivity: largest tau where sensitivity >= target_sensitivity | |
| valid_screening_taus = [sweep_taus[i] for i, s in enumerate(sens_list) if s >= target_sensitivity] | |
| screening_thresh = float(max(valid_screening_taus)) if valid_screening_taus else float(sweep_taus[0]) | |
| modes = { | |
| "default": 0.5, | |
| "youden": youden_thresh, | |
| "f1_optimal": f1_thresh, | |
| "high_sensitivity": screening_thresh, | |
| } | |
| results = {} | |
| for mode_name, tau in modes.items(): | |
| m = compute_metrics(labels, predictions, threshold=tau) | |
| m["threshold"] = tau | |
| results[mode_name] = m | |
| return results | |
| def generate_threshold_sweep( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| thresholds: Optional[np.ndarray] = None, | |
| ) -> pd.DataFrame: | |
| """ | |
| Generate tabular performance sweep across decision thresholds. | |
| """ | |
| if thresholds is None: | |
| thresholds = np.linspace(0.05, 0.95, 19) | |
| rows = [] | |
| for tau in thresholds: | |
| m = compute_metrics(labels, predictions, threshold=float(tau)) | |
| rows.append({ | |
| "Threshold": round(float(tau), 3), | |
| "Accuracy": round(m["accuracy"], 4), | |
| "Balanced_Acc": round(m["balanced_accuracy"], 4), | |
| "Sensitivity": round(m["sensitivity"], 4), | |
| "Specificity": round(m["specificity"], 4), | |
| "PPV": round(m["ppv"], 4), | |
| "NPV": round(m["npv"], 4), | |
| "F1_Score": round(m["f1"], 4), | |
| "MCC": round(m["mcc"], 4), | |
| "Type1_Error": round(m["type1_error_rate"], 4), | |
| "Type2_Error": round(m["type2_error_rate"], 4), | |
| }) | |
| return pd.DataFrame(rows) | |
| def plot_threshold_curves( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| save_path: Optional[str] = None, | |
| target_sensitivity: float = 0.95, | |
| ) -> plt.Figure: | |
| """Plot Sensitivity, Specificity, F1, and Balanced Accuracy across decision thresholds.""" | |
| df_sweep = generate_threshold_sweep(labels, predictions, np.linspace(0.02, 0.98, 100)) | |
| opt = find_optimal_thresholds(labels, predictions, target_sensitivity=target_sensitivity) | |
| fig, ax = plt.subplots(figsize=(9, 6)) | |
| ax.plot(df_sweep["Threshold"], df_sweep["Sensitivity"], label="Sensitivity (Recall)", color="#d9534f", lw=2.2) | |
| ax.plot(df_sweep["Threshold"], df_sweep["Specificity"], label="Specificity", color="#0275d8", lw=2.2) | |
| ax.plot(df_sweep["Threshold"], df_sweep["Balanced_Acc"], label="Balanced Accuracy", color="#5cb85c", lw=2.0, ls="--") | |
| ax.plot(df_sweep["Threshold"], df_sweep["F1_Score"], label="F1-Score", color="#f0ad4e", lw=2.0) | |
| # Vertical lines for optimal operating points | |
| ax.axvline(opt["youden"]["threshold"], color="purple", ls=":", lw=1.8, | |
| label=f"Youden's J ({opt['youden']['threshold']:.2f})") | |
| ax.axvline(opt["high_sensitivity"]["threshold"], color="crimson", ls="-.", lw=1.8, | |
| label=f"High-Sens >= {target_sensitivity*100:.0f}% ({opt['high_sensitivity']['threshold']:.2f})") | |
| ax.set_xlabel("Decision Threshold (tau)", fontsize=13) | |
| ax.set_ylabel("Metric Value", fontsize=13) | |
| ax.set_title("ACL Tear Diagnosis: Threshold-Dependent Performance Curves", fontsize=14, fontweight="bold") | |
| ax.set_xlim(0, 1) | |
| ax.set_ylim(0, 1.02) | |
| ax.grid(alpha=0.3) | |
| ax.legend(loc="lower center", bbox_to_anchor=(0.5, -0.28), ncol=3, fontsize=10) | |
| plt.tight_layout() | |
| if save_path: | |
| os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True) | |
| fig.savefig(save_path, dpi=200, bbox_inches="tight") | |
| return fig | |
| # ββ Dual Confusion Matrix Analysis (Raw + Normalized) βββββββββββββββ | |
| def plot_confusion_matrices_dual( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| threshold: float = 0.5, | |
| save_path: Optional[str] = None, | |
| title: str = "ACL-LKNet Tear Diagnosis: Confusion Matrix Analysis", | |
| ) -> plt.Figure: | |
| """ | |
| Generate side-by-side publication confusion matrix analysis: | |
| Left: Raw integer counts (TN, FP, FN, TP). | |
| Right: Condition-normalized percentages (Sensitivity, Specificity, Type I alpha, Type II beta). | |
| """ | |
| labels = np.asarray(labels).astype(int) | |
| predictions = np.asarray(predictions).astype(float) | |
| tn, fp, fn, tp = _calc_contingency(labels, predictions, threshold) | |
| cm_raw = np.array([[tn, fp], [fn, tp]]) | |
| neg_total = tn + fp if (tn + fp) > 0 else 1 | |
| pos_total = fn + tp if (fn + tp) > 0 else 1 | |
| cm_norm = np.array([ | |
| [tn / neg_total * 100.0, fp / neg_total * 100.0], | |
| [fn / pos_total * 100.0, tp / pos_total * 100.0], | |
| ]) | |
| fig, axes = plt.subplots(1, 2, figsize=(14, 6)) | |
| # 1. Raw Counts Subplot | |
| annot_raw = np.array([ | |
| [f"TN (True Neg)\n{tn}\n({tn/(tn+fp+fn+tp)*100:.1f}%)", f"FP (Type I Error)\n{fp}\n({fp/(tn+fp+fn+tp)*100:.1f}%)"], | |
| [f"FN (Type II Error)\n{fn}\n({fn/(tn+fp+fn+tp)*100:.1f}%)", f"TP (True Pos)\n{tp}\n({tp/(tn+fp+fn+tp)*100:.1f}%)"], | |
| ]) | |
| sns.heatmap( | |
| cm_raw, annot=annot_raw, fmt="", cmap="Blues", cbar=True, ax=axes[0], | |
| xticklabels=["Predicted Negative", "Predicted Positive"], | |
| yticklabels=["Actual Negative", "Actual Positive"], | |
| annot_kws={"size": 11, "weight": "bold"}, | |
| ) | |
| axes[0].set_title(f"A. Absolute Case Counts (tau = {threshold:.2f})", fontsize=13, fontweight="bold") | |
| axes[0].set_xlabel("Predicted Diagnosis", fontsize=12) | |
| axes[0].set_ylabel("True Anatomical State", fontsize=12) | |
| # 2. Condition-Normalized Percentages Subplot | |
| annot_norm = np.array([ | |
| [f"Specificity (TNR)\n{cm_norm[0,0]:.1f}%", f"FPR (Type I alpha)\n{cm_norm[0,1]:.1f}%"], | |
| [f"FNR (Type II beta)\n{cm_norm[1,0]:.1f}%", f"Sensitivity (TPR)\n{cm_norm[1,1]:.1f}%"], | |
| ]) | |
| sns.heatmap( | |
| cm_norm, annot=annot_norm, fmt="", cmap="YlGnBu", cbar=True, ax=axes[1], | |
| vmin=0, vmax=100, | |
| xticklabels=["Predicted Negative", "Predicted Positive"], | |
| yticklabels=["Actual Negative", "Actual Positive"], | |
| annot_kws={"size": 11, "weight": "bold"}, | |
| ) | |
| axes[1].set_title(f"B. Condition-Normalized Diagnostic Rates (%)", fontsize=13, fontweight="bold") | |
| axes[1].set_xlabel("Predicted Diagnosis", fontsize=12) | |
| axes[1].set_ylabel("True Anatomical State", fontsize=12) | |
| fig.suptitle(title, fontsize=15, fontweight="bold", y=1.02) | |
| plt.tight_layout() | |
| if save_path: | |
| os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True) | |
| fig.savefig(save_path, dpi=200, bbox_inches="tight") | |
| return fig | |
| # ββ Quantitative Explainability Evaluation ββββββββββββββββββββββββββ | |
| def evaluate_slice_explainability( | |
| slice_weights: np.ndarray, | |
| target_slice_range: Tuple[int, int] = (10, 18), | |
| labels: Optional[np.ndarray] = None, | |
| ) -> Dict[str, Any]: | |
| """ | |
| Quantify clinical explainability from slice attention distributions: | |
| 1. Pointing Game Hit Rate: % of scans where peak attention falls in anatomical cruciate range. | |
| 2. Cruciate Mass Fraction: Average proportion of attention allocated to central cruciate slices. | |
| 3. Slice Attention Entropy: H(alpha) = -sum alpha_i * log2(alpha_i + eps). (Lower = sharper diagnostic focus). | |
| 4. Hoyer Sparsity: Quantifies degree of attention concentration. | |
| """ | |
| weights = np.asarray(slice_weights).astype(float) | |
| if weights.ndim == 1: | |
| weights = weights.reshape(1, -1) | |
| N, S = weights.shape | |
| row_sums = weights.sum(axis=1, keepdims=True) | |
| row_sums[row_sums == 0] = 1.0 | |
| weights = weights / row_sums | |
| start_s, end_s = target_slice_range | |
| # 1. Pointing Game | |
| peak_slices = np.argmax(weights, axis=1) | |
| hits = (peak_slices >= start_s) & (peak_slices <= end_s) | |
| hit_rate = float(np.mean(hits)) | |
| # 2. Cruciate Attention Mass Fraction | |
| cruciate_mass = weights[:, start_s:end_s + 1].sum(axis=1) | |
| mean_mass_fraction = float(np.mean(cruciate_mass)) | |
| # 3. Attention Entropy | |
| eps = 1e-12 | |
| entropies = -np.sum(weights * np.log2(weights + eps), axis=1) | |
| mean_entropy = float(np.mean(entropies)) | |
| max_entropy = math.log2(S) if S > 1 else 1.0 | |
| norm_entropy = mean_entropy / max_entropy | |
| # 4. Hoyer Sparsity | |
| l1 = np.sum(np.abs(weights), axis=1) | |
| l2 = np.sqrt(np.sum(weights ** 2, axis=1)) | |
| sqrt_s = math.sqrt(S) | |
| sparsities = (sqrt_s - (l1 / (l2 + eps))) / (sqrt_s - 1.0) if sqrt_s > 1 else np.zeros(N) | |
| mean_sparsity = float(np.mean(sparsities)) | |
| summary: Dict[str, Any] = { | |
| "pointing_game_hit_rate": hit_rate, | |
| "cruciate_mass_fraction": mean_mass_fraction, | |
| "attention_entropy_bits": mean_entropy, | |
| "normalized_entropy": norm_entropy, | |
| "hoyer_sparsity": mean_sparsity, | |
| "total_evaluated_cases": int(N), | |
| "target_slice_range": f"Slices {start_s} to {end_s} (of {S})", | |
| } | |
| if labels is not None: | |
| labels = np.asarray(labels).astype(int) | |
| pos_mask = (labels == 1) | |
| neg_mask = (labels == 0) | |
| if np.any(pos_mask): | |
| summary["positive_tear_hit_rate"] = float(np.mean(hits[pos_mask])) | |
| summary["positive_tear_mass_fraction"] = float(np.mean(cruciate_mass[pos_mask])) | |
| summary["positive_tear_entropy"] = float(np.mean(entropies[pos_mask])) | |
| if np.any(neg_mask): | |
| summary["negative_knee_hit_rate"] = float(np.mean(hits[neg_mask])) | |
| summary["negative_knee_mass_fraction"] = float(np.mean(cruciate_mass[neg_mask])) | |
| summary["negative_knee_entropy"] = float(np.mean(entropies[neg_mask])) | |
| return summary | |
| def evaluate_perturbation_faithfulness( | |
| model: nn.Module, | |
| val_loader, | |
| device: torch.device, | |
| k_slices: int = 3, | |
| config=None, | |
| ) -> Dict[str, float]: | |
| """ | |
| Measure quantitative faithfulness of slice attention: | |
| Mask out top-k slices attended by the model and record drop in predicted tear probability. | |
| """ | |
| model.eval() | |
| orig_probs = [] | |
| perturbed_probs = [] | |
| labels = [] | |
| use_amp = config.use_amp if config else True | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| out = model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| p_orig = out["probs"].squeeze(-1).cpu().numpy() | |
| orig_probs.extend(p_orig.tolist()) | |
| labels.extend(batch["label"].numpy().tolist()) | |
| sag_weights = out.get("sag_weights") | |
| if sag_weights is not None: | |
| sag_pert = sag.clone() | |
| B, S = sag_weights.shape | |
| for b in range(B): | |
| top_k_indices = torch.topk(sag_weights[b], k=min(k_slices, S)).indices | |
| sag_pert[b, top_k_indices] = 0.0 | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| out_pert = model(sag_pert, cor, axi, sag_m, cor_m, axi_m) | |
| p_pert = out_pert["probs"].squeeze(-1).cpu().numpy() | |
| perturbed_probs.extend(p_pert.tolist()) | |
| else: | |
| perturbed_probs.extend(p_orig.tolist()) | |
| orig = np.array(orig_probs) | |
| pert = np.array(perturbed_probs) | |
| y = np.array(labels) | |
| delta_p = orig - pert | |
| pos_mask = (y == 1) | |
| return { | |
| "mean_prob_drop_all": float(np.mean(delta_p)), | |
| "mean_prob_drop_positive_cases": float(np.mean(delta_p[pos_mask])) if np.any(pos_mask) else 0.0, | |
| "faithfulness_impact_ratio": float(np.mean(delta_p[pos_mask]) / (np.mean(orig[pos_mask]) + 1e-8)) if np.any(pos_mask) else 0.0, | |
| "top_k_slices_masked": int(k_slices), | |
| } | |
| # ββ Segmentation & Localization Architecture Evaluation βββββββββββββ | |
| def evaluate_slice_localization( | |
| slice_weights: np.ndarray, | |
| ground_truth_active_range: Tuple[int, int] = (10, 18), | |
| threshold: Optional[float] = None, | |
| ) -> Dict[str, float]: | |
| """ | |
| Evaluate slice-level weak volumetric localization using Dice and IoU metrics | |
| against the known central cruciate ligament anatomical zone. | |
| """ | |
| weights = np.asarray(slice_weights).astype(float) | |
| if weights.ndim == 1: | |
| weights = weights.reshape(1, -1) | |
| N, S = weights.shape | |
| gt_mask = np.zeros(S, dtype=float) | |
| start_s, end_s = ground_truth_active_range | |
| gt_mask[start_s:end_s + 1] = 1.0 | |
| soft_dices = [] | |
| hard_dices = [] | |
| ious = [] | |
| for i in range(N): | |
| w = weights[i] | |
| w_norm = (w - w.min()) / (w.max() - w.min() + 1e-8) | |
| # Soft continuous Dice | |
| intersection = np.sum(w_norm * gt_mask) | |
| soft_dice = (2.0 * intersection) / (np.sum(w_norm) + np.sum(gt_mask) + 1e-8) | |
| soft_dices.append(soft_dice) | |
| # Binary Dice & IoU | |
| tau = threshold if threshold is not None else float(np.mean(w_norm)) | |
| pred_bin = (w_norm >= tau).astype(float) | |
| inter_bin = np.sum(pred_bin * gt_mask) | |
| union_bin = np.sum(np.maximum(pred_bin, gt_mask)) | |
| dice_bin = (2.0 * inter_bin) / (np.sum(pred_bin) + np.sum(gt_mask) + 1e-8) | |
| iou_bin = inter_bin / (union_bin + 1e-8) | |
| hard_dices.append(dice_bin) | |
| ious.append(iou_bin) | |
| return { | |
| "slice_localization_soft_dice": float(np.mean(soft_dices)), | |
| "slice_localization_hard_dice": float(np.mean(hard_dices)), | |
| "slice_localization_iou": float(np.mean(ious)), | |
| "anatomical_reference_range": f"Slices {start_s} to {end_s}", | |
| } | |
| def evaluate_msm_reconstruction( | |
| msm_model: nn.Module, | |
| val_loader, | |
| device: torch.device, | |
| ) -> Dict[str, float]: | |
| """ | |
| Evaluate Phase 1 Masked Slice Modeling (MSM) pretext reconstruction fidelity: | |
| Computes MSE and PSNR on masked slice feature representations. | |
| """ | |
| msm_model.eval() | |
| losses = [] | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| out = msm_model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| loss_val = out["loss"].item() | |
| losses.append(loss_val) | |
| mean_mse = float(np.mean(losses)) | |
| psnr = float(10.0 * np.log10(1.0 / (mean_mse + 1e-10))) | |
| return { | |
| "msm_reconstruction_mse": mean_mse, | |
| "msm_reconstruction_psnr_db": psnr, | |
| } | |
| def generate_architecture_benchmark_table(model: nn.Module, config=None) -> pd.DataFrame: | |
| """Generate architectural parameter and memory complexity benchmark table.""" | |
| total_params = sum(p.numel() for p in model.parameters()) | |
| trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) | |
| backbone_params = sum(p.numel() for p in model.backbone.parameters()) if hasattr(model, "backbone") else 0 | |
| slice_attn_params = sum(p.numel() for p in model.slice_attn.parameters()) if hasattr(model, "slice_attn") and model.slice_attn else 0 | |
| fusion_params = sum(p.numel() for p in model.fusion.parameters()) if hasattr(model, "fusion") else 0 | |
| classifier_params = sum(p.numel() for p in model.classifier.parameters()) if hasattr(model, "classifier") else 0 | |
| param_df = pd.DataFrame([ | |
| {"Module / Component": "Shared 2D Backbone (ConvNeXt-Tiny)", "Parameters": f"{backbone_params:,}", "Param Fraction": f"{backbone_params/total_params*100:.1f}%"}, | |
| {"Module / Component": "Slice Attention Pooling", "Parameters": f"{slice_attn_params:,}", "Param Fraction": f"{slice_attn_params/total_params*100:.1f}%"}, | |
| {"Module / Component": "Cross-View Attention Fusion", "Parameters": f"{fusion_params:,}", "Param Fraction": f"{fusion_params/total_params*100:.1f}%"}, | |
| {"Module / Component": "Classification Head", "Parameters": f"{classifier_params:,}", "Param Fraction": f"{classifier_params/total_params*100:.1f}%"}, | |
| {"Module / Component": "TOTAL NETWORK", "Parameters": f"{total_params:,}", "Param Fraction": "100.0%"}, | |
| {"Module / Component": "Trainable Parameters", "Parameters": f"{trainable_params:,}", "Param Fraction": f"{trainable_params/total_params*100:.1f}%"}, | |
| ]) | |
| return param_df | |
| # ββ Cross-Dataset Generalization & Scanner Perturbation Robustness ββ | |
| def evaluate_scanner_perturbation_robustness( | |
| model: nn.Module, | |
| val_loader, | |
| device: torch.device, | |
| config=None, | |
| ) -> Dict[str, Dict[str, float]]: | |
| """ | |
| Stress-test cross-dataset / scanner domain shifts using realistic synthetic perturbations: | |
| 1. Low-Field SNR Shift: Rician noise injection (sigma=0.08) simulating 1.5T scanners. | |
| 2. Slice Thickness Shift: Decimating slice count by 2x and linear interpolating back. | |
| 3. B1 Field Bias Shift: Spatial intensity gradient simulating RF coil inhomogeneity. | |
| """ | |
| model.eval() | |
| use_amp = config.use_amp if config else True | |
| clean_labels = [] | |
| clean_preds = [] | |
| noise_preds = [] | |
| thick_preds = [] | |
| bias_preds = [] | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| clean_labels.extend(batch["label"].numpy().tolist()) | |
| # Baseline | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| out_clean = model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| clean_preds.extend(out_clean["probs"].squeeze(-1).cpu().numpy().tolist()) | |
| # A. Rician Noise Perturbation | |
| noise_sigma = 0.08 | |
| sag_n = torch.sqrt((sag + torch.randn_like(sag)*noise_sigma)**2 + (torch.randn_like(sag)*noise_sigma)**2) | |
| cor_n = torch.sqrt((cor + torch.randn_like(cor)*noise_sigma)**2 + (torch.randn_like(cor)*noise_sigma)**2) | |
| axi_n = torch.sqrt((axi + torch.randn_like(axi)*noise_sigma)**2 + (torch.randn_like(axi)*noise_sigma)**2) | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| out_noise = model(sag_n, cor_n, axi_n, sag_m, cor_m, axi_m) | |
| noise_preds.extend(out_noise["probs"].squeeze(-1).cpu().numpy().tolist()) | |
| # B. Slice Thickness Decimation Perturbation | |
| sag_t = copy.deepcopy(sag) | |
| sag_t[:, 1::2] = sag_t[:, 0::2] | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| out_thick = model(sag_t, cor, axi, sag_m, cor_m, axi_m) | |
| thick_preds.extend(out_thick["probs"].squeeze(-1).cpu().numpy().tolist()) | |
| # C. B1 Field Bias Field | |
| H, W = sag.shape[-2], sag.shape[-1] | |
| y_grad = torch.linspace(0.8, 1.2, H, device=device).unsqueeze(1).repeat(1, W) | |
| sag_b = torch.clamp(sag * y_grad, 0.0, 1.0) | |
| cor_b = torch.clamp(cor * y_grad, 0.0, 1.0) | |
| axi_b = torch.clamp(axi * y_grad, 0.0, 1.0) | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| out_bias = model(sag_b, cor_b, axi_b, sag_m, cor_m, axi_m) | |
| bias_preds.extend(out_bias["probs"].squeeze(-1).cpu().numpy().tolist()) | |
| y = np.array(clean_labels) | |
| base_m = compute_metrics(y, np.array(clean_preds)) | |
| noise_m = compute_metrics(y, np.array(noise_preds)) | |
| thick_m = compute_metrics(y, np.array(thick_preds)) | |
| bias_m = compute_metrics(y, np.array(bias_preds)) | |
| return { | |
| "Baseline": { | |
| "auroc": base_m["auroc"], | |
| "accuracy": base_m["accuracy"], | |
| "sensitivity": base_m["sensitivity"], | |
| "specificity": base_m["specificity"], | |
| "f1": base_m["f1"], | |
| }, | |
| "Low_Field_SNR_Shift (Rician Noise)": { | |
| "auroc": noise_m["auroc"], | |
| "delta_auroc": float(noise_m["auroc"] - base_m["auroc"]), | |
| "accuracy": noise_m["accuracy"], | |
| "delta_accuracy": float(noise_m["accuracy"] - base_m["accuracy"]), | |
| }, | |
| "Slice_Thickness_Decimation": { | |
| "auroc": thick_m["auroc"], | |
| "delta_auroc": float(thick_m["auroc"] - base_m["auroc"]), | |
| "accuracy": thick_m["accuracy"], | |
| "delta_accuracy": float(thick_m["accuracy"] - base_m["accuracy"]), | |
| }, | |
| "B1_Coil_Bias_Field": { | |
| "auroc": bias_m["auroc"], | |
| "delta_auroc": float(bias_m["auroc"] - base_m["auroc"]), | |
| "accuracy": bias_m["accuracy"], | |
| "delta_accuracy": float(bias_m["accuracy"] - base_m["accuracy"]), | |
| }, | |
| } | |
| def evaluate_external_dataset( | |
| model: nn.Module, | |
| external_dataloader, | |
| device: torch.device, | |
| config=None, | |
| dataset_name: str = "External_Knee_Cohort", | |
| ) -> Dict[str, Any]: | |
| """Run out-of-domain evaluation on an external knee MRI dataset.""" | |
| model.eval() | |
| preds = [] | |
| labels = [] | |
| use_amp = config.use_amp if config else True | |
| with torch.no_grad(): | |
| for batch in external_dataloader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| out = model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| preds.extend(out["probs"].squeeze(-1).cpu().numpy().tolist()) | |
| labels.extend(batch["label"].numpy().tolist()) | |
| y = np.array(labels) | |
| p = np.array(preds) | |
| metrics_ci = compute_metrics_with_ci(y, p) | |
| metrics = compute_metrics(y, p) | |
| return { | |
| "dataset_name": dataset_name, | |
| "sample_count": int(len(y)), | |
| "metrics": metrics, | |
| "metrics_ci": metrics_ci, | |
| } | |
| # ββ Statistical Hypothesis Tests βββββββββββββββββββββββββββββββββββ | |
| def mcnemar_test( | |
| labels: np.ndarray, | |
| preds_a: np.ndarray, | |
| preds_b: np.ndarray, | |
| threshold: float = 0.5, | |
| ) -> Tuple[float, float]: | |
| """ | |
| McNemar's test for comparing paired classifier accuracy on the same cohort. | |
| Returns: (chi2_statistic, p_value) | |
| """ | |
| binary_a = (preds_a >= threshold).astype(int) | |
| binary_b = (preds_b >= threshold).astype(int) | |
| correct_a = (binary_a == labels) | |
| correct_b = (binary_b == labels) | |
| b = int(np.sum(correct_a & ~correct_b)) | |
| c = int(np.sum(~correct_a & correct_b)) | |
| if b + c == 0: | |
| return 0.0, 1.0 | |
| chi2 = float((abs(b - c) - 1) ** 2 / (b + c)) | |
| p_value = float(1.0 - stats.chi2.cdf(chi2, df=1)) | |
| return chi2, p_value | |
| def delong_test( | |
| labels: np.ndarray, | |
| preds_a: np.ndarray, | |
| preds_b: np.ndarray, | |
| ) -> Tuple[float, float]: | |
| """ | |
| DeLong's test for comparing AUROCs of two models on paired test set. | |
| Returns: (z_statistic, p_value) | |
| """ | |
| auc_a = float(roc_auc_score(labels, preds_a)) | |
| auc_b = float(roc_auc_score(labels, preds_b)) | |
| n1 = int(np.sum(labels == 1)) | |
| n0 = int(np.sum(labels == 0)) | |
| if n1 == 0 or n0 == 0: | |
| return 0.0, 1.0 | |
| q1_a = auc_a / (2.0 - auc_a) | |
| q2_a = 2.0 * auc_a**2 / (1.0 + auc_a) | |
| se_a = np.sqrt((auc_a * (1.0 - auc_a) + (n1 - 1) * (q1_a - auc_a**2) + | |
| (n0 - 1) * (q2_a - auc_a**2)) / (n1 * n0)) | |
| q1_b = auc_b / (2.0 - auc_b) | |
| q2_b = 2.0 * auc_b**2 / (1.0 + auc_b) | |
| se_b = np.sqrt((auc_b * (1.0 - auc_b) + (n1 - 1) * (q1_b - auc_b**2) + | |
| (n0 - 1) * (q2_b - auc_b**2)) / (n1 * n0)) | |
| se_diff = float(np.sqrt(se_a**2 + se_b**2)) | |
| if se_diff == 0: | |
| return 0.0, 1.0 | |
| z = float((auc_a - auc_b) / se_diff) | |
| p_value = float(2.0 * (1.0 - stats.norm.cdf(abs(z)))) | |
| return z, p_value | |
| # ββ Visualization Helpers ββββββββββββββββββββββββββββββββββββββββββ | |
| def plot_roc_curve( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| save_path: Optional[str] = None, | |
| title: str = "Receiver Operating Characteristic (ROC)", | |
| model_name: str = "ACL-LKNet", | |
| ) -> plt.Figure: | |
| """Plot publication ROC curve.""" | |
| fpr, tpr, _ = roc_curve(labels, predictions) | |
| auc = float(roc_auc_score(labels, predictions)) | |
| fig, ax = plt.subplots(figsize=(7, 7)) | |
| ax.plot(fpr, tpr, color="#0275d8", lw=2.5, label=f"{model_name} (AUROC = {auc:.3f})") | |
| ax.plot([0, 1], [0, 1], color="gray", ls="--", alpha=0.6, label="Chance Reference (AUC = 0.500)") | |
| ax.set_xlabel("False Positive Rate (1 - Specificity)", fontsize=12) | |
| ax.set_ylabel("True Positive Rate (Sensitivity)", fontsize=12) | |
| ax.set_title(title, fontsize=14, fontweight="bold") | |
| ax.legend(loc="lower right", fontsize=11) | |
| ax.grid(alpha=0.3) | |
| plt.tight_layout() | |
| if save_path: | |
| os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True) | |
| fig.savefig(save_path, dpi=200, bbox_inches="tight") | |
| return fig | |
| def plot_precision_recall_curve( | |
| labels: np.ndarray, | |
| predictions: np.ndarray, | |
| save_path: Optional[str] = None, | |
| title: str = "Precision-Recall Curve (AUPRC)", | |
| model_name: str = "ACL-LKNet", | |
| ) -> plt.Figure: | |
| """Plot publication Precision-Recall curve.""" | |
| precision, recall, _ = precision_recall_curve(labels, predictions) | |
| auprc = float(average_precision_score(labels, predictions)) | |
| prevalence = float(labels.sum() / len(labels)) | |
| fig, ax = plt.subplots(figsize=(7, 7)) | |
| ax.plot(recall, precision, color="#5cb85c", lw=2.5, label=f"{model_name} (AUPRC = {auprc:.3f})") | |
| ax.axhline(prevalence, color="gray", ls="--", alpha=0.6, label=f"Prevalence Baseline ({prevalence:.3f})") | |
| ax.set_xlabel("Recall (Sensitivity)", fontsize=12) | |
| ax.set_ylabel("Precision (PPV)", fontsize=12) | |
| ax.set_title(title, fontsize=14, fontweight="bold") | |
| ax.legend(loc="upper right", fontsize=11) | |
| ax.grid(alpha=0.3) | |
| plt.tight_layout() | |
| if save_path: | |
| os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True) | |
| fig.savefig(save_path, dpi=200, bbox_inches="tight") | |
| return fig | |
| def plot_training_curves( | |
| train_history: List[dict], | |
| val_history: List[dict], | |
| save_path: Optional[str] = None, | |
| ) -> plt.Figure: | |
| """Plot multi-panel training and validation dynamics.""" | |
| fig, axes = plt.subplots(1, 3, figsize=(18, 5)) | |
| train_epochs = [h.get("epoch", i) for i, h in enumerate(train_history)] | |
| val_epochs = [h.get("epoch", i) for i, h in enumerate(val_history)] | |
| # Loss | |
| if train_history and "train_loss" in train_history[0]: | |
| axes[0].plot(train_epochs, [h["train_loss"] for h in train_history], label="Train Loss", color="#0275d8", lw=2) | |
| if val_history and "val_loss" in val_history[0]: | |
| axes[0].plot(val_epochs, [h["val_loss"] for h in val_history], label="Val Loss", color="#d9534f", lw=2) | |
| axes[0].set_title("Cross-Entropy / BCE Loss", fontsize=13, fontweight="bold") | |
| axes[0].set_xlabel("Epoch", fontsize=11) | |
| axes[0].legend() | |
| axes[0].grid(alpha=0.3) | |
| # AUROC | |
| if val_history and "val_auroc" in val_history[0]: | |
| axes[1].plot(val_epochs, [h["val_auroc"] for h in val_history], label="Validation AUROC", color="#5cb85c", lw=2) | |
| axes[1].axhline(0.95, color="k", ls="--", alpha=0.5, label="Target (0.95)") | |
| axes[1].set_title("Validation AUROC", fontsize=13, fontweight="bold") | |
| axes[1].set_xlabel("Epoch", fontsize=11) | |
| axes[1].legend() | |
| axes[1].grid(alpha=0.3) | |
| # Accuracy | |
| if val_history and "val_accuracy" in val_history[0]: | |
| axes[2].plot(val_epochs, [h["val_accuracy"] for h in val_history], label="Validation Accuracy", color="#f0ad4e", lw=2) | |
| axes[2].axhline(0.90, color="k", ls="--", alpha=0.5, label="Target (0.90)") | |
| axes[2].set_title("Validation Accuracy", fontsize=13, fontweight="bold") | |
| axes[2].set_xlabel("Epoch", fontsize=11) | |
| axes[2].legend() | |
| axes[2].grid(alpha=0.3) | |
| plt.tight_layout() | |
| if save_path: | |
| os.makedirs(os.path.dirname(os.path.abspath(save_path)), exist_ok=True) | |
| fig.savefig(save_path, dpi=200, bbox_inches="tight") | |
| return fig | |
| def print_results_table(results: Dict[str, Tuple[float, float, float]]): | |
| """Print clean results table with 95% Bootstrap CIs.""" | |
| print("\n" + "=" * 65) | |
| print(f"{'Clinical Metric':<20} {'Point Estimate':>15} {'95% Bootstrap CI':>25}") | |
| print("-" * 65) | |
| for name, (point, lower, upper) in results.items(): | |
| print(f"{name:<20} {point:>15.4f} [{lower:.4f}, {upper:.4f}]") | |
| print("=" * 65 + "\n") | |
| # ββ Master Full Evaluation Orchestrator ββββββββββββββββββββββββββββββ | |
| def full_evaluation( | |
| model: nn.Module, | |
| val_loader, | |
| config, | |
| device: torch.device, | |
| save_dir: Optional[str] = None, | |
| ) -> Dict[str, Any]: | |
| """ | |
| Execute the master scientific evaluation protocol across all 8 reviewer dimensions: | |
| 1. Extended diagnostic classification metrics with 95% Bootstrap CIs. | |
| 2. Thresholding-based evaluations (Youden's J, F1-optimal, high-sensitivity). | |
| 3. Dual confusion matrix generation (counts + condition-normalized). | |
| 4. Quantitative slice explainability (pointing game, entropy, sparsity). | |
| 5. Weak slice localization (Dice, IoU). | |
| 6. Architecture benchmark summary. | |
| 7. Scanner perturbation domain-shift robustness. | |
| """ | |
| model.eval() | |
| all_preds = [] | |
| all_labels = [] | |
| sag_weights_list = [] | |
| use_amp = config.use_amp if config else True | |
| with torch.no_grad(): | |
| for batch in val_loader: | |
| sag = batch["sagittal"].to(device) | |
| cor = batch["coronal"].to(device) | |
| axi = batch["axial"].to(device) | |
| sag_m = batch["sag_mask"].to(device) | |
| cor_m = batch["cor_mask"].to(device) | |
| axi_m = batch["axi_mask"].to(device) | |
| with torch.amp.autocast("cuda", enabled=use_amp and torch.cuda.is_available()): | |
| output = model(sag, cor, axi, sag_m, cor_m, axi_m) | |
| probs = output["probs"].squeeze(-1).cpu().numpy().tolist() | |
| all_preds.extend(probs) | |
| all_labels.extend(batch["label"].numpy().tolist()) | |
| if "sag_weights" in output and output["sag_weights"] is not None: | |
| sag_weights_list.append(output["sag_weights"].cpu().numpy()) | |
| labels = np.array(all_labels) | |
| preds = np.array(all_preds) | |
| # 1. Classification Metrics & Bootstrap CIs | |
| metrics_ci = compute_metrics_with_ci(labels, preds, n_bootstrap=config.bootstrap_n) | |
| print_results_table(metrics_ci) | |
| # 2. Thresholding-based Optimization | |
| target_sens = getattr(config, "target_sensitivity", 0.95) | |
| opt_thresholds = find_optimal_thresholds(labels, preds, target_sensitivity=target_sens) | |
| threshold_sweep_df = generate_threshold_sweep(labels, preds) | |
| chosen_mode = getattr(config, "eval_threshold_mode", "youden") | |
| active_tau = opt_thresholds.get(chosen_mode, opt_thresholds["default"])["threshold"] | |
| # 3. Quantitative Explainability & Localization (if slice weights exist) | |
| explainability_results = {} | |
| localization_results = {} | |
| if len(sag_weights_list) > 0: | |
| all_sag_weights = np.concatenate(sag_weights_list, axis=0) | |
| explain_range = getattr(config, "explainability_target_slice_range", (10, 18)) | |
| explainability_results = evaluate_slice_explainability(all_sag_weights, target_slice_range=explain_range, labels=labels) | |
| localization_results = evaluate_slice_localization(all_sag_weights, ground_truth_active_range=explain_range) | |
| # 4. Scanner Perturbation Robustness | |
| robustness_results = evaluate_scanner_perturbation_robustness(model, val_loader, device, config=config) | |
| # 5. Architecture Benchmark | |
| arch_df = generate_architecture_benchmark_table(model, config=config) | |
| # 6. Save Artifacts & Visualizations | |
| if save_dir: | |
| os.makedirs(save_dir, exist_ok=True) | |
| # Visualizations | |
| plot_roc_curve(labels, preds, os.path.join(save_dir, "roc_curve.png")) | |
| plot_precision_recall_curve(labels, preds, os.path.join(save_dir, "pr_curve.png")) | |
| plot_confusion_matrices_dual(labels, preds, threshold=active_tau, save_path=os.path.join(save_dir, "confusion_matrix_dual.png")) | |
| plot_threshold_curves(labels, preds, save_path=os.path.join(save_dir, "threshold_curves.png"), target_sensitivity=target_sens) | |
| # Tabular data exports | |
| threshold_sweep_df.to_csv(os.path.join(save_dir, "threshold_sweep.csv"), index=False) | |
| arch_df.to_csv(os.path.join(save_dir, "architecture_benchmark.csv"), index=False) | |
| # Save training configuration table if config has export methods | |
| if hasattr(config, "export_config_markdown"): | |
| config.export_config_markdown(os.path.join(save_dir, "training_configuration.md")) | |
| if hasattr(config, "export_config_latex"): | |
| config.export_config_latex(os.path.join(save_dir, "training_configuration.tex")) | |
| # Save summary JSON | |
| summary_payload = { | |
| "metrics": compute_metrics(labels, preds, threshold=active_tau), | |
| "metrics_ci": {k: {"point": float(v[0]), "lower_95": float(v[1]), "upper_95": float(v[2])} for k, v in metrics_ci.items()}, | |
| "optimal_thresholds": opt_thresholds, | |
| "explainability": explainability_results, | |
| "localization": localization_results, | |
| "scanner_robustness": robustness_results, | |
| } | |
| import json | |
| with open(os.path.join(save_dir, "metrics.json"), "w") as f: | |
| json.dump(summary_payload, f, indent=2) | |
| return { | |
| "labels": labels, | |
| "predictions": preds, | |
| "metrics_ci": metrics_ci, | |
| "metrics": compute_metrics(labels, preds, threshold=active_tau), | |
| "optimal_thresholds": opt_thresholds, | |
| "threshold_sweep": threshold_sweep_df, | |
| "explainability": explainability_results, | |
| "localization": localization_results, | |
| "robustness": robustness_results, | |
| "architecture_benchmark": arch_df, | |
| } | |