ACL-LKNet / src /evaluate.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw History Blame Contribute Delete
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,
}