""" 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, }