This commit is contained in:
rpotter6298
2026-07-01 17:35:58 +02:00
parent 9bfcc0243b
commit 35cbd9ac3c
84 changed files with 8500 additions and 423 deletions
+69
View File
@@ -0,0 +1,69 @@
#!/usr/bin/env python3
"""figure_s5.py — Confusion matrix for VGG16 patient-level classification.
Reuses PatientLeakageClassifier.run(return_predictions=True) so the RF ranking,
gamma grid, and final fit are not duplicated here — the same code path that
produces the Figure 4/6 numbers also produces these predictions.
Usage: python scripts/visualizations/figure_s5.py [--tag TAG]
"""
import os
import sys
import argparse
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay
os.environ.setdefault("OMP_NUM_THREADS", "1")
os.environ.setdefault("OPENBLAS_NUM_THREADS", "1")
os.environ.setdefault("MKL_NUM_THREADS", "1")
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(
os.path.abspath(__file__)))))
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
PLOTS_DIR = os.path.join(ROOT, "plots")
from classes import PatientLeakageClassifier
SEED = 20
N_JOBS = 6
CLASS_NAMES = ["Benign", "Malignant", "Normal"]
def display_name(raw):
"""Map a raw dataset label ('Bengin cases', ...) to a display class name."""
if raw.startswith("Bengin"):
return "Benign"
if raw.startswith("Malignant"):
return "Malignant"
return "Normal"
ap = argparse.ArgumentParser()
ap.add_argument("--tag", default="")
args = ap.parse_args()
tag = f"_{args.tag}" if args.tag else ""
clf = PatientLeakageClassifier(
os.path.join(ROOT, "results", "simple_patient_manifest.csv"),
os.path.join(ROOT, "features"), n_jobs=N_JOBS)
r = clf.run("VGG16", SEED, "patient", return_predictions=True)
print(f"Best: n={r['nfeat']}, gamma={r['gamma']:.6e}, "
f"CV={r['cv']:.4f}, Test={r['test']:.4f}")
y_true = [display_name(c) for c in r["y_true"]]
y_pred = [display_name(c) for c in r["y_pred"]]
cm = confusion_matrix(y_true, y_pred, labels=CLASS_NAMES, normalize="true")
fig, ax = plt.subplots(figsize=(6, 5))
ConfusionMatrixDisplay(cm, display_labels=CLASS_NAMES).plot(
cmap="Blues", ax=ax, colorbar=True, values_format=".2f")
ax.set_title("Figure S5 — Patient-level Confusion Matrix (VGG16)", fontsize=12)
plt.tight_layout()
out = os.path.join(PLOTS_DIR, "figure_s5", f"figure_s5{tag}.png")
os.makedirs(os.path.dirname(out), exist_ok=True)
plt.savefig(out, dpi=150)
plt.close()
print(f"Saved → {out}")