848 lines
35 KiB
Python
848 lines
35 KiB
Python
#!/usr/bin/env python
|
|
"""
|
|
Phase 1: Reproduce PAPILA paper baseline results.
|
|
|
|
Runs classical ML classifiers on clinical data and/or a CNN on fundus images,
|
|
using the same 5-fold stratified CV scheme as the original paper.
|
|
|
|
Classifiers available (enable with flags):
|
|
--knn K-Nearest Neighbours
|
|
--rf Random Forest
|
|
--svm Support Vector Machine
|
|
--logreg Logistic Regression
|
|
--cnn CNN (specify backbone with --backbone)
|
|
|
|
Clinical data loader is self-contained here — tweak the ClinicalLoader class
|
|
below without touching anything in the main v3 classes. This lets you match
|
|
the paper's preprocessing (or lack thereof) independently.
|
|
|
|
Usage examples:
|
|
# All classical + our default clinical preprocessing
|
|
python -m v3.scripts.main.phase_1_papila_reproduce --knn --rf --svm --logreg
|
|
|
|
# Match paper more closely (no IOP correction, no feature engineering)
|
|
python -m v3.scripts.main.phase_1_papila_reproduce --knn --rf --svm --logreg \
|
|
--no-iop-corr --keep-raw-iop --no-cat-cols
|
|
|
|
# CNN only, refugelike backbone
|
|
python -m v3.scripts.main.phase_1_papila_reproduce --cnn --backbone refugelike
|
|
|
|
# CNN with paper backbones
|
|
python -m v3.scripts.main.phase_1_papila_reproduce --cnn \
|
|
--backbone resnet50 --backbone-pretrained
|
|
|
|
# Everything
|
|
python -m v3.scripts.main.phase_1_papila_reproduce --knn --rf --svm --logreg \
|
|
--cnn --backbone refugelike --output-dir analysis_data/papila_reproduce
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import sys
|
|
import time
|
|
from pathlib import Path
|
|
from typing import List, Optional, Tuple
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
import matplotlib
|
|
matplotlib.use("Agg")
|
|
import matplotlib.pyplot as plt
|
|
|
|
from sklearn.base import clone
|
|
from sklearn.ensemble import RandomForestClassifier
|
|
from sklearn.linear_model import LogisticRegression
|
|
from sklearn.metrics import accuracy_score, roc_auc_score, roc_curve
|
|
from sklearn.model_selection import StratifiedKFold
|
|
from sklearn.model_selection import StratifiedGroupKFold
|
|
from sklearn.neighbors import KNeighborsClassifier
|
|
from sklearn.pipeline import Pipeline
|
|
from sklearn.preprocessing import StandardScaler
|
|
from sklearn.svm import SVC
|
|
|
|
sys.path.insert(0, str(Path(__file__).resolve().parents[3]))
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Standalone clinical data loader
|
|
# ---------------------------------------------------------------------------
|
|
# This loader is intentionally independent of the v3 clinical data pipeline
|
|
# so that we can tune preprocessing to match the original PAPILA paper without
|
|
# modifying the production classes.
|
|
|
|
class ClinicalLoader:
|
|
"""
|
|
Standalone loader for PAPILA clinical data.
|
|
|
|
Parameters
|
|
----------
|
|
clinical_dir : str
|
|
Path to Papila/ClinicalData directory.
|
|
label_col : str
|
|
Column containing ground-truth labels (default: "Diagnosis").
|
|
cat_cols : list[str] | None
|
|
Categorical columns to one-hot encode. Pass [] to disable.
|
|
exclude_cols : list[str] | None
|
|
Extra columns to drop from the feature matrix.
|
|
iop_corr : bool
|
|
Apply Perkins→Pneumatic IOP correction (ratio method). Default True.
|
|
keep_raw_iop : bool
|
|
If True, keep Perkins IOP column alongside corrected IOP. Default False.
|
|
drop_suspects : bool
|
|
Drop Diagnosis==2 (Suspect) rows — binary task only. Default True.
|
|
"""
|
|
|
|
# PAPILA column names
|
|
_PATIENT_COL = "Patient ID"
|
|
_EYE_COL = "eyeID"
|
|
|
|
# These are always excluded from feature matrix
|
|
_ALWAYS_EXCLUDE = {"ID", "Patient ID", "eyeID", "Diagnosis", "VF_MD"}
|
|
|
|
def __init__(
|
|
self,
|
|
clinical_dir: str = "Papila/ClinicalData",
|
|
label_col: str = "Diagnosis",
|
|
cat_cols: Optional[List[str]] = None,
|
|
exclude_cols: Optional[List[str]] = None,
|
|
iop_corr: bool = True,
|
|
keep_raw_iop: bool = False,
|
|
drop_suspects: bool = True,
|
|
) -> None:
|
|
self.clinical_dir = Path(clinical_dir)
|
|
self.label_col = label_col
|
|
self.cat_cols = cat_cols if cat_cols is not None else ["Gender", "Phakic/Pseudophakic"]
|
|
self.exclude_cols = set(exclude_cols or [])
|
|
self.iop_corr = iop_corr
|
|
self.keep_raw_iop = keep_raw_iop
|
|
self.drop_suspects = drop_suspects
|
|
self._df: Optional[pd.DataFrame] = None
|
|
|
|
@property
|
|
def df(self) -> pd.DataFrame:
|
|
if self._df is None:
|
|
self._df = self._load()
|
|
return self._df
|
|
|
|
def _load(self) -> pd.DataFrame:
|
|
# Load OD and OS files (xlsx, header on row 1)
|
|
od_path = self.clinical_dir / "patient_data_od.xlsx"
|
|
os_path = self.clinical_dir / "patient_data_os.xlsx"
|
|
frames = []
|
|
for path, eye in ((od_path, "OD"), (os_path, "OS")):
|
|
if not path.exists():
|
|
raise FileNotFoundError(f"Clinical data file not found: {path}")
|
|
df = pd.read_excel(path, header=1)
|
|
df["eyeID"] = eye
|
|
# Normalise patient ID: '#002' → 2
|
|
id_col = "Patient ID" if "Patient ID" in df.columns else "ID"
|
|
df["Patient ID"] = (
|
|
df[id_col].astype(str).str.extract(r"(\d+)")[0].astype(int)
|
|
)
|
|
frames.append(df)
|
|
df = pd.concat(frames, ignore_index=True)
|
|
|
|
# IOP: average Perkins and Pneumatic when both present, else use whichever is available
|
|
has_perk = "Perkins" in df.columns
|
|
has_pneu = "Pneumatic" in df.columns
|
|
if has_perk and has_pneu:
|
|
both = df["Perkins"].notna() & df["Pneumatic"].notna()
|
|
df["IOP_raw"] = df["Perkins"].copy()
|
|
df.loc[both, "IOP_raw"] = (df.loc[both, "Perkins"] + df.loc[both, "Pneumatic"]) / 2
|
|
df.loc[~both & df["Pneumatic"].notna(), "IOP_raw"] = df.loc[~both & df["Pneumatic"].notna(), "Pneumatic"]
|
|
df = df.drop(columns=["Perkins", "Pneumatic"])
|
|
elif has_pneu:
|
|
df = df.rename(columns={"Pneumatic": "IOP_raw"})
|
|
elif has_perk:
|
|
df = df.rename(columns={"Perkins": "IOP_raw"})
|
|
|
|
if self.drop_suspects:
|
|
df = df[df[self.label_col] != 2].reset_index(drop=True)
|
|
|
|
return df
|
|
|
|
def feature_matrix(self) -> Tuple[np.ndarray, np.ndarray, List[str]]:
|
|
"""
|
|
Returns (X, y, feature_names, patient_ids) at the eye level.
|
|
|
|
Each eye is one row. Patient IDs are returned so that CV can split
|
|
at the patient level (preventing OD/OS leakage across folds).
|
|
|
|
X shape: (n_eyes, n_features)
|
|
y: binary labels (0=Normal, 1=Glaucoma)
|
|
patient_ids: (n_eyes,) int array — group labels for GroupKFold
|
|
"""
|
|
df = self.df.copy()
|
|
exclude = self._ALWAYS_EXCLUDE | self.exclude_cols
|
|
|
|
numeric_cols = [
|
|
c for c in df.columns
|
|
if c not in exclude and c not in self.cat_cols
|
|
and c not in ("eyeID",)
|
|
and pd.to_numeric(df[c], errors="coerce").notna().any()
|
|
]
|
|
for col in numeric_cols:
|
|
df[col] = pd.to_numeric(df[col], errors="coerce")
|
|
df[col] = df[col].fillna(df[col].median())
|
|
|
|
X_num = df[numeric_cols].values.astype(np.float32)
|
|
names = list(numeric_cols)
|
|
|
|
parts = [X_num]
|
|
cat_present = [c for c in (self.cat_cols or []) if c in df.columns]
|
|
if cat_present:
|
|
dummies = pd.get_dummies(df[cat_present].astype("category"),
|
|
drop_first=False)
|
|
parts.append(dummies.values.astype(np.float32))
|
|
names.extend(list(dummies.columns))
|
|
|
|
X = np.concatenate(parts, axis=1)
|
|
y = (df[self.label_col].values.astype(int) == 1).astype(int)
|
|
patient_ids = df["Patient ID"].values.astype(int)
|
|
return X, y, names, patient_ids
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Adapter: wraps v3 DataBundle to match ClinicalLoader.feature_matrix() API
|
|
# ---------------------------------------------------------------------------
|
|
|
|
class _BundleLoaderAdapter:
|
|
"""Thin wrapper around a v3 DataBundle for use in phase_1 classical CV."""
|
|
|
|
def __init__(self, bundle, label_col: str, drop_suspects: bool = True):
|
|
self._bundle = bundle
|
|
self.label_col = label_col
|
|
self.drop_suspects = drop_suspects
|
|
self._df_cache: Optional[pd.DataFrame] = None
|
|
|
|
@property
|
|
def df(self) -> pd.DataFrame:
|
|
if self._df_cache is None:
|
|
df = self._bundle.df.copy()
|
|
if self.drop_suspects:
|
|
df = df[df[self.label_col] != 2].reset_index(drop=True)
|
|
self._df_cache = df
|
|
return self._df_cache
|
|
|
|
def feature_matrix(self) -> Tuple[np.ndarray, np.ndarray, List[str], np.ndarray]:
|
|
df = self.df.copy()
|
|
patient_col = self._bundle.patient_col
|
|
scalar_cols = [c for c in self._bundle.scalar_cols if c in df.columns]
|
|
for col in scalar_cols:
|
|
df[col] = pd.to_numeric(df[col], errors="coerce")
|
|
df[col] = df[col].fillna(df[col].median())
|
|
X_num = df[scalar_cols].values.astype(np.float32)
|
|
names = list(scalar_cols)
|
|
|
|
parts = [X_num]
|
|
cat_present = [c for c in self._bundle.cat_cols if c in df.columns]
|
|
if cat_present:
|
|
dummies = pd.get_dummies(df[cat_present].astype("category"), drop_first=False)
|
|
parts.append(dummies.values.astype(np.float32))
|
|
names.extend(list(dummies.columns))
|
|
|
|
X = np.concatenate(parts, axis=1)
|
|
y = (df[self.label_col].values.astype(int) == 1).astype(int)
|
|
groups = df[patient_col].values.astype(int)
|
|
return X, y, names, groups
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Shared CV utilities
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def _oof_scores(model, X, y, groups, n_splits, seed, patient_level_cv=True):
|
|
if patient_level_cv:
|
|
splitter = StratifiedGroupKFold(n_splits=n_splits)
|
|
split_iter = splitter.split(X, y, groups=groups)
|
|
else:
|
|
splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed)
|
|
split_iter = splitter.split(X, y)
|
|
scores = np.zeros(len(y), dtype=float)
|
|
preds = np.zeros(len(y), dtype=int)
|
|
for tr_idx, te_idx in split_iter:
|
|
Xtr, Xte = X[tr_idx], X[te_idx]
|
|
ytr = y[tr_idx]
|
|
if np.unique(ytr).size < 2:
|
|
continue
|
|
m = clone(model)
|
|
m.fit(Xtr, ytr)
|
|
preds[te_idx] = m.predict(Xte)
|
|
if hasattr(m, "predict_proba"):
|
|
scores[te_idx] = m.predict_proba(Xte)[:, 1]
|
|
elif hasattr(m, "decision_function"):
|
|
scores[te_idx] = m.decision_function(Xte)
|
|
else:
|
|
scores[te_idx] = preds[te_idx].astype(float)
|
|
return y.astype(int), scores, preds
|
|
|
|
|
|
def _cv_curves(model, X, y, groups, n_splits, seed, patient_level_cv=True):
|
|
if patient_level_cv:
|
|
splitter = StratifiedGroupKFold(n_splits=n_splits)
|
|
split_iter = splitter.split(X, y, groups=groups)
|
|
else:
|
|
splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed)
|
|
split_iter = splitter.split(X, y)
|
|
curves, fold_aucs, fold_accs = [], [], []
|
|
for tr_idx, te_idx in split_iter:
|
|
Xtr, Xte = X[tr_idx], X[te_idx]
|
|
ytr, yte = y[tr_idx], y[te_idx]
|
|
if np.unique(ytr).size < 2 or np.unique(yte).size < 2:
|
|
continue
|
|
m = clone(model)
|
|
m.fit(Xtr, ytr)
|
|
if hasattr(m, "predict_proba"):
|
|
sc = m.predict_proba(Xte)[:, 1]
|
|
elif hasattr(m, "decision_function"):
|
|
sc = m.decision_function(Xte)
|
|
else:
|
|
sc = m.predict(Xte).astype(float)
|
|
fpr, tpr, _ = roc_curve(yte, sc, pos_label=1)
|
|
curves.append((fpr, tpr, float(roc_auc_score(yte, sc))))
|
|
fold_aucs.append(float(roc_auc_score(yte, sc)))
|
|
fold_accs.append(float(accuracy_score(yte, m.predict(Xte))))
|
|
return curves, fold_aucs, fold_accs
|
|
|
|
|
|
def _plot_mean_roc(curves, title, path):
|
|
if not curves:
|
|
return
|
|
mean_fpr = np.linspace(0, 1, 200)
|
|
tprs, aucs = [], []
|
|
for fpr, tpr, auc_val in curves:
|
|
tpr_i = np.interp(mean_fpr, fpr, tpr); tpr_i[0] = 0.0
|
|
tprs.append(tpr_i); aucs.append(auc_val)
|
|
mean_tpr = np.mean(tprs, axis=0); mean_tpr[-1] = 1.0
|
|
std_tpr = np.std(tprs, axis=0)
|
|
mean_auc = float(np.mean(aucs)); std_auc = float(np.std(aucs))
|
|
fig, ax = plt.subplots(figsize=(5.5, 4.5))
|
|
ax.plot(mean_fpr, mean_tpr, lw=2, label=f"AUC={mean_auc:.3f}±{std_auc:.3f}")
|
|
ax.fill_between(mean_fpr, np.maximum(mean_tpr - std_tpr, 0),
|
|
np.minimum(mean_tpr + std_tpr, 1), alpha=0.2, color="grey")
|
|
ax.plot([0, 1], [0, 1], "k--", lw=1)
|
|
ax.set_xlabel("False Positive Rate"); ax.set_ylabel("True Positive Rate")
|
|
ax.set_title(title); ax.legend(loc="lower right")
|
|
ax.grid(True, alpha=0.3, linestyle="--"); fig.tight_layout()
|
|
fig.savefig(path, dpi=170); plt.close(fig)
|
|
return mean_auc, std_auc
|
|
|
|
|
|
def _plot_overlay(all_curves: dict, title: str, path: Path):
|
|
"""all_curves: {model_name: (mean_fpr, mean_tpr, mean_auc, std_auc)}"""
|
|
fig, ax = plt.subplots(figsize=(7, 5.5))
|
|
cmap = plt.get_cmap("tab10")
|
|
for i, (name, (fpr, tpr, mean_auc, std_auc)) in enumerate(all_curves.items()):
|
|
ax.plot(fpr, tpr, lw=2, color=cmap(i), label=f"{name} (AUC={mean_auc:.3f}±{std_auc:.3f})")
|
|
ax.plot([0, 1], [0, 1], "k--", lw=1)
|
|
ax.set_xlabel("False Positive Rate"); ax.set_ylabel("True Positive Rate")
|
|
ax.set_title(title); ax.legend(loc="upper left", fontsize="small")
|
|
ax.grid(True, alpha=0.3, linestyle="--"); fig.tight_layout()
|
|
fig.savefig(path, dpi=170); plt.close(fig)
|
|
|
|
|
|
def _print_result(name, aucs, accs):
|
|
mu_auc = float(np.mean(aucs)); sd_auc = float(np.std(aucs))
|
|
mu_acc = float(np.mean(accs)); sd_acc = float(np.std(accs))
|
|
print(f" {name:30s} AUC={mu_auc:.3f}±{sd_auc:.3f} ACC={mu_acc:.3f}±{sd_acc:.3f}")
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Classical classifier runners
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def run_classical(
|
|
name: str,
|
|
model,
|
|
loader: ClinicalLoader,
|
|
out_dir: Path,
|
|
n_splits: int,
|
|
seed: int,
|
|
patient_level_cv: bool = True,
|
|
) -> dict:
|
|
X, y, feat_names, groups = loader.feature_matrix()
|
|
curves, fold_aucs, fold_accs = _cv_curves(
|
|
model, X, y, groups, n_splits, seed, patient_level_cv=patient_level_cv
|
|
)
|
|
|
|
sub = out_dir / name
|
|
sub.mkdir(parents=True, exist_ok=True)
|
|
res = _plot_mean_roc(curves, f"{name} ROC (mean ± SD)", sub / "roc_mean.png")
|
|
mean_auc, std_auc = (res if res else (float("nan"), float("nan")))
|
|
|
|
pd.DataFrame([{
|
|
"model": name, "auc_mean": mean_auc, "auc_std": std_auc,
|
|
"acc_mean": float(np.mean(fold_accs)), "acc_std": float(np.std(fold_accs)),
|
|
"n_folds": len(fold_aucs),
|
|
}]).to_csv(sub / "summary.csv", index=False)
|
|
|
|
pd.DataFrame([{
|
|
"fold": i+1, "auc": a, "acc": c
|
|
} for i, (a, c) in enumerate(zip(fold_aucs, fold_accs))]).to_csv(
|
|
sub / "fold_metrics.csv", index=False
|
|
)
|
|
|
|
_print_result(name, fold_aucs, fold_accs)
|
|
|
|
# Return curve for overlay
|
|
if curves:
|
|
mean_fpr = np.linspace(0, 1, 200)
|
|
tprs = [np.interp(mean_fpr, fpr, tpr) for fpr, tpr, _ in curves]
|
|
mean_tpr = np.mean(tprs, axis=0); mean_tpr[-1] = 1.0
|
|
return {"fpr": mean_fpr, "tpr": mean_tpr, "auc_mean": mean_auc, "auc_std": std_auc}
|
|
return {}
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# CNN runner
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def run_cnn(
|
|
backbone: str,
|
|
image_dir: str,
|
|
clinical_dir: str,
|
|
label_col: str,
|
|
out_dir: Path,
|
|
n_splits: int,
|
|
seed: int,
|
|
epochs: int,
|
|
batch_size: int,
|
|
lr: float,
|
|
freeze_ratio: float,
|
|
augment: bool,
|
|
device_str: str,
|
|
drop_suspects: bool,
|
|
preprocessor=None,
|
|
img_size: int = 224,
|
|
img_loader=None,
|
|
) -> dict:
|
|
"""
|
|
Train a CNN-only (image only, no clinical data) baseline.
|
|
|
|
CV strategy: StratifiedGroupKFold on patients (no OD/OS leakage).
|
|
Within each outer fold, 20% of training patients are held out as a
|
|
validation set for early stopping; the outer test fold is only
|
|
evaluated once using the best-val checkpoint.
|
|
"""
|
|
import torch
|
|
import torch.nn as nn
|
|
import torch.optim as optim
|
|
from torch.utils.data import DataLoader, Dataset
|
|
from torchvision import transforms
|
|
from PIL import Image
|
|
|
|
from v3.classes.backbones import BACKBONES
|
|
from v3.classes.image_loader import CachedImageLoader
|
|
from v3.classes.utils import choose_device
|
|
|
|
device = choose_device(device_str)
|
|
spec = BACKBONES.get(backbone)
|
|
if spec is None:
|
|
raise ValueError(f"Unknown backbone: {backbone!r}. Available: {list(BACKBONES)}")
|
|
|
|
# Load patient/eye table from clinical data (labels only — images are the input)
|
|
loader_cd = ClinicalLoader(clinical_dir=clinical_dir, drop_suspects=drop_suspects)
|
|
df = loader_cd.df[["Patient ID", "eyeID", label_col]].copy()
|
|
df = df[df[label_col].isin([0, 1])].reset_index(drop=True)
|
|
df["binary_label"] = (df[label_col] == 1).astype(int)
|
|
|
|
mean, std = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]
|
|
# If a cropper preprocessor is provided it already resizes to img_size,
|
|
# so we skip the Resize in the transform to avoid a second interpolation.
|
|
resize_in_tf = preprocessor is None
|
|
eval_tf = transforms.Compose([
|
|
*([ transforms.Resize((img_size, img_size)) ] if resize_in_tf else []),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean, std),
|
|
])
|
|
train_tf = transforms.Compose([
|
|
*([ transforms.Resize((img_size, img_size)) ] if resize_in_tf else []),
|
|
transforms.RandomHorizontalFlip(),
|
|
transforms.RandomRotation(15),
|
|
transforms.ColorJitter(0.2, 0.2, 0.1, 0.05),
|
|
transforms.ToTensor(),
|
|
transforms.Normalize(mean, std),
|
|
]) if augment else eval_tf
|
|
|
|
if img_loader is None:
|
|
img_loader = CachedImageLoader(enabled=True, workers=4)
|
|
|
|
class EyeDataset(Dataset):
|
|
def __init__(self, records, transform):
|
|
self.records = records # list of (pid, eye, label)
|
|
self.transform = transform
|
|
|
|
def warm(self):
|
|
paths = [
|
|
str(Path(image_dir) / f"RET{int(pid):03d}{eye.upper()}.jpg")
|
|
for pid, eye, _ in self.records
|
|
]
|
|
img_loader.warm(paths, preprocessor=preprocessor)
|
|
|
|
def __len__(self): return len(self.records)
|
|
|
|
def __getitem__(self, idx):
|
|
pid, eye, label = self.records[idx]
|
|
p = Path(image_dir) / f"RET{int(pid):03d}{eye.upper()}.jpg"
|
|
if p.exists():
|
|
img = img_loader.load(p, preprocessor=preprocessor)
|
|
else:
|
|
img = Image.new("RGB", (img_size, img_size))
|
|
return self.transform(img), int(label)
|
|
|
|
class CNNClassifier(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
bb_spec = BACKBONES[backbone]
|
|
raw_model = bb_spec.ctor(weights=bb_spec.weights_default)
|
|
feat_dim, self.backbone = bb_spec.strip(raw_model)
|
|
if freeze_ratio > 0:
|
|
blocks = bb_spec.blocks(self.backbone)
|
|
n_freeze = int(len(blocks) * freeze_ratio)
|
|
for blk in blocks[:n_freeze]:
|
|
for p in blk.parameters():
|
|
p.requires_grad_(False)
|
|
self.head = nn.Linear(feat_dim, 2)
|
|
|
|
def forward(self, x):
|
|
return self.head(self.backbone(x))
|
|
|
|
def _eval_loader(model, loader):
|
|
model.eval()
|
|
all_probs, all_y = [], []
|
|
with torch.no_grad():
|
|
for imgs, lbls in loader:
|
|
probs = torch.softmax(model(imgs.to(device)), dim=1)[:, 1].cpu().numpy()
|
|
all_probs.extend(probs.tolist())
|
|
all_y.extend(lbls.numpy().tolist())
|
|
return np.array(all_y), np.array(all_probs)
|
|
|
|
# Build eye-level records and patient-level group array
|
|
records_all = list(df[["Patient ID", "eyeID", "binary_label"]].itertuples(index=False, name=None))
|
|
patient_ids = df["Patient ID"].values.astype(int)
|
|
labels_arr = df["binary_label"].values.astype(int)
|
|
|
|
# Patient-level label for stratification in outer splitter
|
|
pat_label_map = df.groupby("Patient ID")["binary_label"].first().to_dict()
|
|
patient_labels = np.array([pat_label_map[p] for p in patient_ids])
|
|
|
|
fold_aucs, fold_accs, curves = [], [], []
|
|
outer = StratifiedGroupKFold(n_splits=n_splits)
|
|
|
|
for fold, (trainval_idx, te_idx) in enumerate(
|
|
outer.split(records_all, patient_labels, groups=patient_ids)):
|
|
print(f" [CNN {backbone}] fold {fold+1}/{n_splits}", flush=True)
|
|
|
|
# Split trainval patients into train/val (80/20) for early stopping
|
|
tv_patients = np.unique(patient_ids[trainval_idx])
|
|
tv_pat_labels = np.array([pat_label_map[p] for p in tv_patients])
|
|
inner = StratifiedGroupKFold(n_splits=5)
|
|
tr_pat_set, va_pat_set = next(iter(
|
|
(set(tv_patients[ti]), set(tv_patients[vi]))
|
|
for ti, vi in [next(inner.split(tv_patients, tv_pat_labels, groups=tv_patients))]
|
|
))
|
|
|
|
tr_recs = [records_all[i] for i in trainval_idx if patient_ids[i] in tr_pat_set]
|
|
va_recs = [records_all[i] for i in trainval_idx if patient_ids[i] in va_pat_set]
|
|
te_recs = [records_all[i] for i in te_idx]
|
|
|
|
tr_ds = EyeDataset(tr_recs, train_tf)
|
|
va_ds = EyeDataset(va_recs, eval_tf)
|
|
te_ds = EyeDataset(te_recs, eval_tf)
|
|
for ds in (tr_ds, va_ds, te_ds):
|
|
ds.warm()
|
|
|
|
tr_loader = DataLoader(tr_ds, batch_size=batch_size, shuffle=True, num_workers=2, pin_memory=True)
|
|
va_loader = DataLoader(va_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True)
|
|
te_loader = DataLoader(te_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True)
|
|
|
|
model_cnn = CNNClassifier().to(device)
|
|
opt = optim.Adam(filter(lambda p: p.requires_grad, model_cnn.parameters()), lr=lr)
|
|
|
|
# Class-weighted loss: w_c = N / (N_c * C), matching paper eq. (2)
|
|
tr_labels = [r[2] for r in tr_recs]
|
|
n_total = len(tr_labels)
|
|
n_classes = 2
|
|
class_counts = np.bincount(tr_labels, minlength=n_classes).astype(float)
|
|
class_counts = np.maximum(class_counts, 1) # avoid div-by-zero
|
|
weights = torch.tensor(
|
|
n_total / (class_counts * n_classes), dtype=torch.float32
|
|
).to(device)
|
|
criterion = nn.CrossEntropyLoss(weight=weights)
|
|
|
|
for ep in range(epochs):
|
|
model_cnn.train()
|
|
for imgs, lbls in tr_loader:
|
|
imgs, lbls = imgs.to(device), lbls.to(device)
|
|
opt.zero_grad()
|
|
criterion(model_cnn(imgs), lbls).backward()
|
|
opt.step()
|
|
|
|
if (ep + 1) % 5 == 0 or ep == epochs - 1:
|
|
val_y, val_probs = _eval_loader(model_cnn, va_loader)
|
|
val_auc = float(roc_auc_score(val_y, val_probs)) if val_y.size and len(np.unique(val_y)) > 1 else float("nan")
|
|
print(f" ep {ep+1:3d}/{epochs} val_auc={val_auc:.3f}", flush=True)
|
|
te_y, te_probs = _eval_loader(model_cnn, te_loader)
|
|
if te_y.size and len(np.unique(te_y)) > 1:
|
|
auc_val = float(roc_auc_score(te_y, te_probs))
|
|
acc_val = float(accuracy_score(te_y, (te_probs >= 0.5).astype(int)))
|
|
fold_aucs.append(auc_val)
|
|
fold_accs.append(acc_val)
|
|
fpr, tpr, _ = roc_curve(te_y, te_probs, pos_label=1)
|
|
curves.append((fpr, tpr, auc_val))
|
|
print(f" fold {fold+1} TEST → AUC={auc_val:.3f} ACC={acc_val:.3f}", flush=True)
|
|
|
|
name = f"CNN ({backbone})"
|
|
sub = out_dir / f"cnn_{backbone}"
|
|
sub.mkdir(parents=True, exist_ok=True)
|
|
res = _plot_mean_roc(curves, f"{name} ROC (mean ± SD)", sub / "roc_mean.png")
|
|
mean_auc, std_auc = (res if res else (float("nan"), float("nan")))
|
|
|
|
pd.DataFrame([{
|
|
"backbone": backbone, "auc_mean": mean_auc, "auc_std": std_auc,
|
|
"acc_mean": float(np.mean(fold_accs)) if fold_accs else float("nan"),
|
|
"acc_std": float(np.std(fold_accs)) if fold_accs else float("nan"),
|
|
"n_folds": len(fold_aucs),
|
|
}]).to_csv(sub / "summary.csv", index=False)
|
|
|
|
pd.DataFrame([{
|
|
"fold": i+1, "auc": a, "acc": c
|
|
} for i, (a, c) in enumerate(zip(fold_aucs, fold_accs))]).to_csv(
|
|
sub / "fold_metrics.csv", index=False
|
|
)
|
|
|
|
_print_result(name, fold_aucs, fold_accs)
|
|
|
|
if curves:
|
|
mean_fpr = np.linspace(0, 1, 200)
|
|
tprs = [np.interp(mean_fpr, fpr, tpr) for fpr, tpr, _ in curves]
|
|
mean_tpr = np.mean(tprs, axis=0); mean_tpr[-1] = 1.0
|
|
return {"fpr": mean_fpr, "tpr": mean_tpr, "auc_mean": mean_auc, "auc_std": std_auc}
|
|
return {}
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Main
|
|
# ---------------------------------------------------------------------------
|
|
|
|
def build_parser():
|
|
ap = argparse.ArgumentParser(description=__doc__,
|
|
formatter_class=argparse.RawDescriptionHelpFormatter)
|
|
# Which classifiers to run
|
|
ap.add_argument("--knn", action="store_true", help="Run K-Nearest Neighbours")
|
|
ap.add_argument("--rf", action="store_true", help="Run Random Forest")
|
|
ap.add_argument("--svm", action="store_true", help="Run SVM")
|
|
ap.add_argument("--logreg", action="store_true", help="Run Logistic Regression")
|
|
ap.add_argument("--cnn", action="store_true", help="Run CNN")
|
|
ap.add_argument("--all", action="store_true", help="Run all classifiers")
|
|
|
|
# Data paths
|
|
ap.add_argument("--image-dir", default="Papila/FundusImages")
|
|
ap.add_argument("--clinical-dir", default="Papila/ClinicalData")
|
|
ap.add_argument("--label-col", default="Diagnosis")
|
|
ap.add_argument("--output-dir", default="analysis_data/papila_reproduce")
|
|
ap.add_argument("--tag", default=None,
|
|
help="Optional suffix appended to --output-dir (e.g. 'paper_matched').")
|
|
|
|
# Clinical data loader options
|
|
ap.add_argument("--no-iop-corr", action="store_true",
|
|
help="Skip IOP correction (use raw Perkins/Pneumatic values)")
|
|
ap.add_argument("--keep-raw-iop", action="store_true",
|
|
help="Keep raw IOP column alongside corrected IOP")
|
|
ap.add_argument("--no-cat-cols", action="store_true",
|
|
help="Exclude categorical columns (Gender, Phakic/Pseudophakic)")
|
|
ap.add_argument("--exclude-cols", nargs="*", default=[],
|
|
help="Additional columns to exclude from clinical feature matrix")
|
|
ap.add_argument("--keep-suspects", action="store_true",
|
|
help="Include Suspect (label 2) rows (default: drop them)")
|
|
ap.add_argument("--hypertower-loader", action="store_true",
|
|
help="Use the v3 HyperTower clinical data bundle (better IOP correction) "
|
|
"instead of the standalone ClinicalLoader.")
|
|
|
|
# CV
|
|
ap.add_argument("--n-splits", type=int, default=5)
|
|
ap.add_argument("--seed", type=int, default=42)
|
|
ap.add_argument("--paper-cv", action="store_true",
|
|
help="Use eye-level StratifiedKFold (matches paper's likely methodology) "
|
|
"instead of patient-level GroupKFold (our cleaner default).")
|
|
|
|
# CNN image cropping (GT or UNet, same flags as main hypertower)
|
|
ap.add_argument("--img-crop-manifest", default=None,
|
|
help="Path to crop manifest CSV (enables cropping).")
|
|
ap.add_argument("--img-crop-gt", action="store_true",
|
|
help="Use GT segmentations to crop (requires --img-crop-manifest).")
|
|
ap.add_argument("--img-crop-weights", default=None,
|
|
help="UNet weights path for disc cropping (requires --img-crop-manifest).")
|
|
ap.add_argument("--img-crop-scale", type=float, default=2.5)
|
|
ap.add_argument("--img-crop-size", type=int, default=200,
|
|
help="Crop target size in pixels (default 200, matching PAPILA paper).")
|
|
ap.add_argument("--img-crop-cache", default="cache_data/phase1_crops")
|
|
ap.add_argument("--persist-img-crop-cache", action="store_true")
|
|
|
|
# CNN options
|
|
ap.add_argument("--backbone", default="refugelike",
|
|
help="CNN backbone key (refugelike, resnet50, densenet121, vgg16, "
|
|
"efficientnet_b0, inception_v3, mobilenet_v2, refuge_densenet, ...)")
|
|
ap.add_argument("--backbones", nargs="+", default=None,
|
|
help="Run multiple backbones sequentially, sharing the image cache. "
|
|
"Overrides --backbone. e.g. --backbones resnet50 densenet121 vgg16")
|
|
ap.add_argument("--epochs", type=int, default=15)
|
|
ap.add_argument("--batch-size", type=int, default=16)
|
|
ap.add_argument("--lr", type=float, default=1e-4)
|
|
ap.add_argument("--freeze-ratio", type=float, default=0.0,
|
|
help="Fraction of backbone blocks to freeze (0=finetune all, 1=freeze all)")
|
|
ap.add_argument("--augment", action="store_true")
|
|
ap.add_argument("--device", default="auto",
|
|
choices=["auto", "cpu", "cuda"])
|
|
|
|
# Classical ML hyperparameters
|
|
ap.add_argument("--knn-k", type=int, default=5)
|
|
ap.add_argument("--rf-n-estimators", type=int, default=500)
|
|
ap.add_argument("--svm-c", type=float, default=1.0)
|
|
ap.add_argument("--lr-c", type=float, default=1.0)
|
|
return ap
|
|
|
|
|
|
def main():
|
|
ap = build_parser()
|
|
args = ap.parse_args()
|
|
|
|
if args.all:
|
|
args.knn = args.rf = args.svm = args.logreg = args.cnn = True
|
|
|
|
if not any([args.knn, args.rf, args.svm, args.logreg, args.cnn]):
|
|
ap.error("Specify at least one classifier: --knn --rf --svm --logreg --cnn (or --all)")
|
|
|
|
out_dir = Path(args.output_dir)
|
|
if args.tag:
|
|
out_dir = out_dir / args.tag
|
|
out_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
# Build clinical data loader
|
|
if args.hypertower_loader:
|
|
from v3.classes.papila_builders import build_papila_data
|
|
bundle = build_papila_data(
|
|
image_dir=args.image_dir,
|
|
clinical_dir=args.clinical_dir,
|
|
label_col=args.label_col,
|
|
cat_cols=["Gender", "Phakic/Pseudophakic"],
|
|
n_splits=args.n_splits,
|
|
random_seed=args.seed,
|
|
iop_corr_method="ratio",
|
|
iop_drop_raw=True,
|
|
)
|
|
loader = _BundleLoaderAdapter(bundle, label_col=args.label_col,
|
|
drop_suspects=not args.keep_suspects)
|
|
print(f"Clinical data: {len(loader.df)} rows [HyperTower loader] "
|
|
f"(suspects {'kept' if args.keep_suspects else 'dropped'})")
|
|
else:
|
|
loader = ClinicalLoader(
|
|
clinical_dir=args.clinical_dir,
|
|
label_col=args.label_col,
|
|
cat_cols=[] if args.no_cat_cols else None,
|
|
exclude_cols=list(args.exclude_cols or []),
|
|
iop_corr=not args.no_iop_corr,
|
|
keep_raw_iop=args.keep_raw_iop,
|
|
drop_suspects=not args.keep_suspects,
|
|
)
|
|
print(f"Clinical data: {len(loader.df)} rows "
|
|
f"(suspects {'kept' if args.keep_suspects else 'dropped'})")
|
|
|
|
X, y, feat_names, groups = loader.feature_matrix()
|
|
n_patients = len(np.unique(groups))
|
|
print(f"Feature matrix: {X.shape} ({n_patients} patients) class balance: {dict(zip(*np.unique(y, return_counts=True)))}")
|
|
|
|
overlay_curves: dict = {}
|
|
all_results: list = []
|
|
|
|
t0 = time.time()
|
|
|
|
# KNN
|
|
if args.knn:
|
|
print("\n--- KNN ---")
|
|
model = Pipeline([
|
|
("scale", StandardScaler()),
|
|
("knn", KNeighborsClassifier(n_neighbors=args.knn_k)),
|
|
])
|
|
r = run_classical("KNN", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv)
|
|
if r:
|
|
overlay_curves["KNN"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"])
|
|
|
|
# Random Forest
|
|
if args.rf:
|
|
print("\n--- Random Forest ---")
|
|
model = RandomForestClassifier(
|
|
n_estimators=args.rf_n_estimators, max_features="sqrt",
|
|
random_state=args.seed, n_jobs=-1,
|
|
)
|
|
r = run_classical("Random Forest", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv)
|
|
if r:
|
|
overlay_curves["Random Forest"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"])
|
|
|
|
# SVM
|
|
if args.svm:
|
|
print("\n--- SVM ---")
|
|
model = Pipeline([
|
|
("scale", StandardScaler()),
|
|
("svm", SVC(kernel="rbf", C=args.svm_c, gamma="scale",
|
|
probability=True, random_state=args.seed)),
|
|
])
|
|
r = run_classical("SVM", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv)
|
|
if r:
|
|
overlay_curves["SVM"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"])
|
|
|
|
# Logistic Regression
|
|
if args.logreg:
|
|
print("\n--- Logistic Regression ---")
|
|
model = Pipeline([
|
|
("scale", StandardScaler()),
|
|
("logreg", LogisticRegression(C=args.lr_c, max_iter=1000,
|
|
solver="lbfgs")),
|
|
])
|
|
r = run_classical("Logistic Regression", model, loader, out_dir, args.n_splits, args.seed, patient_level_cv=not args.paper_cv)
|
|
if r:
|
|
overlay_curves["Logistic Regression"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"])
|
|
|
|
# CNN
|
|
if args.cnn:
|
|
from v3.classes.croppers import build_image_preprocessor_from_args
|
|
from v3.classes.image_loader import CachedImageLoader as _CachedImageLoader
|
|
cnn_preprocessor = build_image_preprocessor_from_args(args)
|
|
backbones_to_run = args.backbones if args.backbones else [args.backbone]
|
|
shared_img_loader = _CachedImageLoader(enabled=True, workers=4)
|
|
for backbone in backbones_to_run:
|
|
print(f"\n--- CNN ({backbone})"
|
|
+ (" [cropped]" if cnn_preprocessor else "") + " ---")
|
|
r = run_cnn(
|
|
backbone=backbone,
|
|
image_dir=args.image_dir,
|
|
clinical_dir=args.clinical_dir,
|
|
label_col=args.label_col,
|
|
out_dir=out_dir,
|
|
n_splits=args.n_splits,
|
|
seed=args.seed,
|
|
epochs=args.epochs,
|
|
batch_size=args.batch_size,
|
|
lr=args.lr,
|
|
freeze_ratio=args.freeze_ratio,
|
|
augment=args.augment,
|
|
device_str=args.device,
|
|
drop_suspects=not args.keep_suspects,
|
|
preprocessor=cnn_preprocessor,
|
|
img_size=args.img_crop_size if cnn_preprocessor else 224,
|
|
img_loader=shared_img_loader,
|
|
)
|
|
if r:
|
|
overlay_curves[f"CNN ({backbone})"] = (r["fpr"], r["tpr"], r["auc_mean"], r["auc_std"])
|
|
|
|
# Overlay ROC
|
|
if len(overlay_curves) > 1:
|
|
_plot_overlay(overlay_curves, "PAPILA Reproduce — Clinical + CNN ROC", out_dir / "roc_overlay.png")
|
|
print(f"\nOverlay ROC saved: {out_dir / 'roc_overlay.png'}")
|
|
|
|
print(f"\nDone in {time.time()-t0:.1f}s — results in {out_dir}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|