update 3-19

This commit is contained in:
rpotter6298
2026-03-19 11:18:58 +01:00
parent 7ea85d5426
commit 786457b30d
35 changed files with 4019 additions and 258 deletions
@@ -0,0 +1,399 @@
#!/usr/bin/env python3
"""v2 port of cnn_logits_rf_cv: CNN logit extraction + RF CV using the v2 PAPILA stack.
Replaces the hardcoded-config v1 version. Data loading, fold splitting, and
feature preparation all go through the v2 stack so that --exclude-cols,
--iop-corr-method, etc. are first-class options.
Usage:
python scripts/basic_analysis/cnn_logits_rf_cv_v2.py \
--eval-mode binary --backbone resnet50 --epochs 40 \
--exclude-cols Phakic/Pseudophakic Axial_Length \
--run-name cnn_rf_nocrop_binary
"""
from __future__ import annotations
import argparse
import random
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import sys
from types import SimpleNamespace
from typing import Dict, List, Tuple
import numpy as np
import pandas as pd
from PIL import Image
import torch
from torch import nn
from torch.utils.data import DataLoader
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score, roc_auc_score
REPO_ROOT = Path(__file__).resolve().parents[2]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from classes.v2.papila_builders import build_papila_data
from classes.v2.backbones import BACKBONES, load_backbone_weights
from classes.v2.dataset import ClinicalDataset, _ClinicalView
from classes.v2.split_manager import PatientFirstSplitManager
from classes.v2.transforms import build_backbone_transform, build_eval_transform
# ---------------------------------------------------------------------------
# Model
# ---------------------------------------------------------------------------
class CNNHead(nn.Module):
def __init__(self, backbone_name: str, num_classes: int) -> None:
super().__init__()
spec = BACKBONES[backbone_name]
backbone = spec.ctor(weights=spec.weights_default)
if backbone_name.startswith("refuge"):
load_backbone_weights(backbone_name, backbone)
out_dim, backbone = spec.strip(backbone)
self.backbone = backbone
self.head = nn.Linear(out_dim, num_classes)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.head(self.backbone(x))
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _auc_score(y_true: np.ndarray, probs: np.ndarray, num_classes: int) -> float:
try:
if num_classes == 2:
return float(roc_auc_score(y_true, probs[:, 1]))
return float(roc_auc_score(y_true, probs, multi_class="ovr", average="macro"))
except Exception:
return float("nan")
def _train_cnn(
model: nn.Module,
loader: DataLoader,
args,
fold: int,
device: torch.device,
) -> None:
model.train()
optimizer = torch.optim.Adam(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
criterion = nn.CrossEntropyLoss()
for epoch in range(args.epochs):
running_loss = total = correct = 0
for x_img, _meta, y in loader:
x_img = x_img.to(device)
y = y.to(device=device, dtype=torch.long)
optimizer.zero_grad()
logits = model(x_img)
loss = criterion(logits, y)
loss.backward()
optimizer.step()
running_loss += float(loss.item()) * int(y.size(0))
correct += int((logits.argmax(1) == y).sum().item())
total += int(y.size(0))
if (epoch + 1) % args.log_every == 0:
print(
f" [fold {fold+1}] epoch {epoch+1}/{args.epochs} "
f"loss={running_loss/max(total,1):.4f} acc={correct/max(total,1):.4f}",
flush=True,
)
def _infer_logits(
model: nn.Module, loader: DataLoader, device: torch.device
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Returns (y_true, logits, metadata_vectors)."""
model.eval()
y_all, logits_all, md_all = [], [], []
with torch.no_grad():
for x_img, x_md, y in loader:
logits = model(x_img.to(device)).cpu().numpy()
y_np = y.numpy() if torch.is_tensor(y) else np.asarray(y)
md_np = x_md.numpy() if torch.is_tensor(x_md) else np.asarray(x_md)
y_all.append(y_np)
logits_all.append(logits)
md_all.append(md_np)
return (
np.concatenate(y_all),
np.concatenate(logits_all),
np.concatenate(md_all),
)
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def build_parser() -> argparse.ArgumentParser:
ap = argparse.ArgumentParser(
description="CNN logit extraction + RF CV (v2 PAPILA stack)"
)
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("--cat-cols", nargs="*", default=["Gender", "Phakic/Pseudophakic"])
ap.add_argument("--exclude-cols", nargs="*", default=[],
help="Feature columns to drop from the clinical feature matrix.")
ap.add_argument("--eval-mode", choices=["binary", "multiclass"], default="binary")
ap.add_argument("--n-splits", type=int, default=5)
ap.add_argument("--fold-seed", type=int, default=42)
ap.add_argument("--holdout-per-class", type=int, default=6,
help="Patients per class reserved for holdout (0 disables).")
ap.add_argument("--holdout-seed", type=int, default=123)
ap.add_argument("--backbone", default="resnet50")
ap.add_argument("--batch-size", type=int, default=8)
ap.add_argument("--epochs", type=int, default=40)
ap.add_argument("--lr", type=float, default=1e-4)
ap.add_argument("--weight-decay", type=float, default=1e-5)
ap.add_argument("--rf-trees", type=int, default=500)
ap.add_argument("--rf-max-depth", type=int, default=None)
ap.add_argument("--rf-min-samples-leaf", type=int, default=1)
ap.add_argument("--num-workers", type=int, default=0)
ap.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto")
ap.add_argument("--seed", type=int, default=1234)
ap.add_argument("--log-every", type=int, default=5)
ap.add_argument("--run-name", default=None)
ap.add_argument("--output-root", default="analysis_data/basic_analysis/cnn_logits_rf_cv_v2")
ap.add_argument("--iop-corr-method", choices=["ratio", "ols", "lad", "multi"], default="ratio")
ap.add_argument("--iop-drop-raw", action="store_true", default=False)
ap.add_argument("--in-memory-cache", action="store_true", default=True,
help="Cache all images in RAM before training (default: on).")
ap.add_argument("--no-in-memory-cache", action="store_false", dest="in_memory_cache",
help="Disable in-memory image cache.")
ap.add_argument("--cache-workers", type=int, default=4,
help="Threads for prebuilding image cache (default: 4).")
return ap
def _prebuild_image_cache(df: "pd.DataFrame", data, n_workers: int,
resize: int = 256) -> dict:
"""Load, convert to RGB, resize to `resize`px short edge, and cache as uint8 arrays.
Storing pre-resized images means the per-batch transform only has to do
CenterCrop + augmentation + ToTensor + Normalize on a small image rather
than resizing a full-resolution fundus image every step.
"""
paths = list({str(data.get_image_path(row)) for _, row in df.iterrows()})
cache: dict = {}
print(f"[cache] Preloading {len(paths)} images (resize={resize}px) "
f"with {n_workers} threads...", flush=True)
def _load(p: str):
img = Image.open(p)
if img.mode != "RGB":
img = img.convert("RGB")
w, h = img.size
scale = resize / min(w, h)
img = img.resize((round(w * scale), round(h * scale)), Image.BILINEAR)
return p, np.asarray(img, dtype=np.uint8)
with ThreadPoolExecutor(max_workers=max(1, n_workers)) as ex:
for path, arr in ex.map(_load, paths):
cache[path] = arr
ex_shape = cache[paths[0]].shape
print(f"[cache] Done — {len(cache)} images in RAM "
f"({ex_shape[1]}×{ex_shape[0]} each).", flush=True)
return cache
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main() -> None:
args = build_parser().parse_args()
random.seed(args.seed)
np.random.seed(args.seed)
torch.manual_seed(args.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(args.seed)
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
print(f"Device: {device}", flush=True)
run_name = args.run_name or time.strftime("%Y%m%d_%H%M%S")
out_dir = Path(args.output_root) / run_name
out_dir.mkdir(parents=True, exist_ok=True)
# ---------- data ----------
print("Loading PAPILA data...", flush=True)
data = build_papila_data(
image_dir=args.image_dir,
clinical_dir=args.clinical_dir,
label_col=args.label_col,
cat_cols=list(args.cat_cols),
n_splits=args.n_splits,
random_seed=args.fold_seed,
iop_corr_method=args.iop_corr_method,
iop_drop_raw=args.iop_drop_raw,
exclude_cols=list(args.exclude_cols or []),
)
print(f"Loaded: {len(data.df)} rows feature_dim={data.feature_dim}", flush=True)
df_mode = data.df.copy()
if args.eval_mode == "binary":
df_mode = df_mode[df_mode[args.label_col].isin([0, 1])].reset_index(drop=True)
num_classes = 2 if args.eval_mode == "binary" else int(df_mode[args.label_col].nunique())
print(
f"eval_mode={args.eval_mode} num_classes={num_classes} "
f"rows={len(df_mode)} patients={df_mode['Patient ID'].nunique()}",
flush=True,
)
# ---------- splits ----------
split_manager = PatientFirstSplitManager(
patient_col="Patient ID", label_col=args.label_col
)
split_args = SimpleNamespace(
eval_mode=args.eval_mode,
holdout_per_class=args.holdout_per_class,
holdout_seed=args.holdout_seed,
n_splits=args.n_splits,
fold_seed=args.fold_seed,
)
clinical_ns = SimpleNamespace(df=df_mode, label_col=args.label_col)
plans = split_manager.build_plans(clinical=clinical_ns, args=split_args, profile=None)
n_folds = min(args.n_splits, len(plans))
if plans and plans[0].holdout is not None:
plans[0].holdout.to_csv(out_dir / "holdout_patients.csv", index=False)
# ---------- image cache ----------
image_cache = None
if args.in_memory_cache:
image_cache = _prebuild_image_cache(df_mode, data, args.cache_workers)
# ---------- transforms ----------
train_tf = build_backbone_transform(args.backbone, augment=True)
eval_tf = build_eval_transform(args.backbone)
rows: List[Dict] = []
holdout_rows: List[Dict] = []
for fold, split in enumerate(plans[:n_folds]):
print(f"\n[info] Fold {fold+1}/{n_folds}", flush=True)
random.seed(args.seed + fold * 100)
np.random.seed(args.seed + fold * 100)
torch.manual_seed(args.seed + fold * 100)
view_train = _ClinicalView(data, split.train)
view_val = _ClinicalView(data, split.val)
dl_train = DataLoader(
ClinicalDataset(view_train, train_tf, image_cache=image_cache),
batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers,
)
dl_val = DataLoader(
ClinicalDataset(view_val, eval_tf, image_cache=image_cache),
batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers,
)
dl_holdout = None
if split.holdout is not None and not split.holdout.empty:
dl_holdout = DataLoader(
ClinicalDataset(_ClinicalView(data, split.holdout), eval_tf,
image_cache=image_cache),
batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers,
)
# Train CNN
model = CNNHead(args.backbone, num_classes=num_classes).to(device)
_train_cnn(model, dl_train, args, fold=fold, device=device)
# Extract logits (re-run train without augmentation for RF features)
dl_train_eval = DataLoader(
ClinicalDataset(view_train, eval_tf, image_cache=image_cache),
batch_size=args.batch_size, shuffle=False, num_workers=args.num_workers,
)
y_tr, log_tr, md_tr = _infer_logits(model, dl_train_eval, device)
y_va, log_va, md_va = _infer_logits(model, dl_val, device)
np.save(out_dir / f"fold{fold}_train_logits.npy", log_tr)
np.save(out_dir / f"fold{fold}_val_logits.npy", log_va)
X_tr = np.concatenate([log_tr, md_tr], axis=1)
X_va = np.concatenate([log_va, md_va], axis=1)
rf = RandomForestClassifier(
n_estimators=args.rf_trees,
max_depth=args.rf_max_depth,
min_samples_leaf=args.rf_min_samples_leaf,
class_weight="balanced",
random_state=args.fold_seed + fold,
n_jobs=-1,
)
rf.fit(X_tr, y_tr)
p_va = rf.predict_proba(X_va)
pred_va = np.argmax(p_va, axis=1)
rows.append({
"fold": fold,
"val_acc": float(accuracy_score(y_va, pred_va)),
"val_auc": _auc_score(y_va, p_va, num_classes),
"n_val": int(len(y_va)),
})
if dl_holdout is not None:
y_ho, log_ho, md_ho = _infer_logits(model, dl_holdout, device)
np.save(out_dir / f"fold{fold}_holdout_logits.npy", log_ho)
X_ho = np.concatenate([log_ho, md_ho], axis=1)
p_ho = rf.predict_proba(X_ho)
pred_ho = np.argmax(p_ho, axis=1)
holdout_rows.append({
"fold": fold,
"holdout_acc": float(accuracy_score(y_ho, pred_ho)),
"holdout_auc": _auc_score(y_ho, p_ho, num_classes),
"n_holdout": int(len(y_ho)),
})
print(
f"[info] Fold {fold+1} RF: val_acc={rows[-1]['val_acc']:.4f}"
f" val_auc={rows[-1]['val_auc']:.4f}"
f" | holdout_acc={holdout_rows[-1]['holdout_acc']:.4f}"
f" holdout_auc={holdout_rows[-1]['holdout_auc']:.4f}",
flush=True,
)
else:
print(
f"[info] Fold {fold+1} RF: val_acc={rows[-1]['val_acc']:.4f}"
f" val_auc={rows[-1]['val_auc']:.4f}",
flush=True,
)
# ---------- save + summarise ----------
fold_df = pd.DataFrame(rows)
fold_df.to_csv(out_dir / "rf_val_metrics.csv", index=False)
print("\nRF validation metrics:")
print(fold_df.to_string(index=False, float_format=lambda x: f"{x:.4f}"))
if holdout_rows:
ho_df = pd.DataFrame(holdout_rows)
ho_df.to_csv(out_dir / "rf_holdout_metrics.csv", index=False)
print("\nRF holdout metrics:")
print(ho_df.to_string(index=False, float_format=lambda x: f"{x:.4f}"))
print(
f"\nMeans: val_acc={fold_df['val_acc'].mean():.4f}"
f" val_auc={fold_df['val_auc'].mean():.4f}"
f" holdout_acc={ho_df['holdout_acc'].mean():.4f}"
f" holdout_auc={ho_df['holdout_auc'].mean():.4f}"
)
else:
print(
f"\nMeans: val_acc={fold_df['val_acc'].mean():.4f}"
f" val_auc={fold_df['val_auc'].mean():.4f}"
)
print(f"\nSaved outputs to: {out_dir}", flush=True)
if __name__ == "__main__":
main()