Files
hypertower/v3/classes/v3_hypertower.py
T
2026-03-19 16:58:29 +01:00

1251 lines
64 KiB
Python

"""V3HyperTower — central orchestrator for the V3 pipeline.
Key differences from V2:
- Proper outer/inner k-fold CV: test = current fold, val = next fold, train = rest.
No pre-carved holdout — every patient appears in test exactly once.
- No checkpoint saving (.pt files). Model states are kept in memory only.
- Test set evaluated each main-phase epoch (logged to epoch CSV only, never used for model selection).
Final test metrics reported in summary use the last-epoch model state.
- Binary-focused defaults (multiclass still supported via --eval-mode multiclass).
- tune_binary_threshold uses Youden's J by default (class-distribution independent).
- tune_multiclass_bias maximises balanced accuracy (class-distribution independent).
"""
from __future__ import annotations
import argparse
import csv
import json
import time
from pathlib import Path
from types import SimpleNamespace
from typing import Optional
import numpy as np
import torch
from v3.classes.croppers import build_image_preprocessor_from_args
from v3.classes.image_loader import CachedImageLoader
from v3.classes.dataset import _ClinicalView # noqa: F401
from v3.classes.loader_factory import (
build_balanced_sampler,
filter_bilateral_samples,
filter_eye_samples,
make_loader,
)
from v3.classes.metrics import _score_arrays, _svf, _tune_and_snap
from v3.classes.models import (
BilateralHT,
FusedEnsembleHT,
SingleEyeHT,
V2ModeComparisonOps,
collect_probs_bilateral,
collect_probs_bilateral_components,
collect_probs_classic,
collect_probs_ensemble,
collect_probs_ensemble_pereye,
collect_probs_eye_level,
collect_probs_fused,
collect_probs_single_components,
train_bilateral_epoch,
train_fusion_epoch,
train_single_epoch,
)
from v3.classes.papila_builders import build_papila_data
from v3.classes.predictions import PredictionStore, head_names_for_mode
from v3.classes.profiles import build_papila_profile
from v3.classes.results import FoldArtifacts, FoldResult, _f, _nan, _sv
from v3.classes.split_manager import PatientFirstSplitManager
from v3.classes.transforms import build_eval_transform
from v3.classes.utils import (
_drop_mixed_label_patients,
_relabel_mixed_patients_to_max,
choose_device,
seed_everything,
)
from v3.classes.hypertower_logger import HypertowerLogger
# ---------------------------------------------------------------------------
# Helpers (unchanged from V2)
# ---------------------------------------------------------------------------
def _fusion_events(y, pf, pi, pm):
pred_f = pf.argmax(1); pred_i = pi.argmax(1); pred_m = pm.argmax(1)
corr = int(((pred_f == y) & (pred_i != y) & (pred_m != y)).sum())
err = int(((pred_f != y) & (pred_i == y) & (pred_m == y)).sum())
return corr, err
def _cm_cells(y, p, num_classes):
if not y.size or p.ndim < 2 or p.shape[1] != num_classes:
return {}
pred = p.argmax(1)
if num_classes == 2:
return {
"tn": int(((pred==0)&(y==0)).sum()), "fp": int(((pred==1)&(y==0)).sum()),
"fn": int(((pred==0)&(y==1)).sum()), "tp": int(((pred==1)&(y==1)).sum()),
}
out = {}
for i in range(num_classes):
for j in range(num_classes):
out[f"cm_{i}_{j}"] = int(((y==i)&(pred==j)).sum())
return out
def _save_predictions_csv(fold_dir, eval_mode, y_true, heads, suffix=""):
N = len(y_true)
num_classes = next(p.shape[1] for p in heads.values() if p is not None)
rows = []
for i in range(N):
true = int(y_true[i])
row = {"idx": i, "y_true": true}
for head_name, probs in heads.items():
if probs is None:
continue
pred = int(probs[i].argmax())
row[f"pred_{head_name}"] = pred
for c in range(num_classes):
row[f"prob_{head_name}_c{c}"] = float(probs[i, c])
if eval_mode == "binary":
row[f"tp_{head_name}"] = int(pred==1 and true==1)
row[f"fp_{head_name}"] = int(pred==1 and true==0)
row[f"tn_{head_name}"] = int(pred==0 and true==0)
row[f"fn_{head_name}"] = int(pred==0 and true==1)
else:
row[f"correct_{head_name}"] = int(pred==true)
rows.append(row)
if not rows:
return
path = fold_dir / f"predictions{suffix}.csv"
with path.open("w", newline="") as f:
writer = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
writer.writeheader()
writer.writerows(rows)
# ---------------------------------------------------------------------------
# V3HyperTower
# ---------------------------------------------------------------------------
class V3HyperTower:
"""V3 orchestrator. Construct with ``V3HyperTower(args)``, call ``.run()``."""
@staticmethod
def build_parser() -> argparse.ArgumentParser:
ap = argparse.ArgumentParser(
description=(
"V3 HyperTower — outer/inner k-fold CV (test=current fold). "
"No pre-carved holdout. No checkpoint saving."
)
)
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=[])
ap.add_argument("--eval-mode", choices=["binary", "multiclass"], default="binary")
ap.add_argument(
"--tower-mode", choices=["single", "ensemble", "bilateral", "classic"],
default="ensemble",
)
ap.add_argument("--n-splits", type=int, default=5)
ap.add_argument("--fold-seed", type=int, default=42)
ap.add_argument(
"--folds", type=int, default=None,
help="Optional cap on number of folds to run.",
)
ap.add_argument("--epochs", type=int, default=40)
ap.add_argument("--warmup-tower-epochs", type=int, default=None)
ap.add_argument("--warmup-fused-epochs", type=int, default=None)
ap.add_argument("--single-warmup-tower-epochs", type=int, default=None)
ap.add_argument("--single-warmup-fused-epochs", type=int, default=None)
ap.add_argument("--warmup-cd-epochs", type=int, default=0)
ap.add_argument("--bilat-warmup-tower-epochs", type=int, default=None)
ap.add_argument("--bilat-warmup-fused-epochs", type=int, default=None)
ap.add_argument("--batch-size", type=int, default=8)
ap.add_argument("--lr", type=float, default=1e-4)
ap.add_argument("--bcd-prob", type=float, default=0.5)
ap.add_argument("--tower-loss-mode", choices=["bcd", "all"], default="bcd")
ap.add_argument("--backbone", default="refugelike")
ap.add_argument("--freeze-ratio", type=float, default=0.0)
ap.add_argument("--augment", action="store_true")
ap.add_argument("--balanced-sampling", action="store_true")
ap.add_argument("--num-workers", type=int, default=4)
ap.add_argument("--in-memory-cache", action="store_true", default=True)
ap.add_argument("--no-in-memory-cache", action="store_false", dest="in_memory_cache")
ap.add_argument("--cache-workers", type=int, default=4)
ap.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto")
ap.add_argument("--seed", type=int, default=1234)
ap.add_argument("--run-name", default=None)
ap.add_argument("--output-root", default="analysis_data")
# ROI cropping
ap.add_argument("--img-crop-manifest", type=str, default=None)
ap.add_argument("--img-crop-gt", action="store_true")
ap.add_argument("--img-crop-weights", type=str, default=None)
ap.add_argument("--img-crop-normalize", type=str, default="per_image",
choices=["per_image", "imagenet"])
ap.add_argument("--img-crop-threshold", type=float, default=0.5)
ap.add_argument("--img-crop-tta", action="store_true")
ap.add_argument("--img-crop-scale", type=float, default=2.5)
ap.add_argument("--img-crop-size", type=int, default=224)
ap.add_argument("--img-crop-cache", type=str, default="cache_data/hypertower_crops")
ap.add_argument("--persist-img-crop-cache", action="store_true")
# Architecture
ap.add_argument("--cd-hidden-dim", type=int, default=128)
ap.add_argument("--fusion-dim", type=int, default=256)
ap.add_argument("--bridge-mode", default="fused",
choices=["fused", "image_only", "clinical_only"])
# Mixed patients
ap.add_argument("--exclude-mixed-patients", dest="exclude_mixed_patients",
action="store_true")
ap.add_argument("--include-mixed-patients", dest="exclude_mixed_patients",
action="store_false")
ap.add_argument("--relabel-mixed-patients-to-max", dest="relabel_mixed_patients_to_max",
action="store_true")
ap.add_argument("--keep-mixed-raw-labels", dest="relabel_mixed_patients_to_max",
action="store_false", help=argparse.SUPPRESS)
ap.set_defaults(exclude_mixed_patients=False, relabel_mixed_patients_to_max=False)
# Tuning
ap.add_argument("--tune-binary-threshold", action="store_true")
ap.add_argument("--tune-multiclass-bias", action="store_true")
ap.add_argument("--ece-bins", type=int, default=10)
ap.add_argument("--log-every", type=int, default=1)
# IOP
ap.add_argument("--iop-corr-method", choices=["ratio", "ols", "lad", "multi"],
default="ratio")
ap.add_argument("--iop-drop-raw", action="store_true", default=False)
# Fused head
ap.add_argument("--fused-head", action="store_true")
ap.add_argument("--fusion-epochs", type=int, default=10)
return ap
def __init__(self, args) -> None:
self.args = args
self.device = choose_device(args.device)
seed_everything(args.seed)
print(f"Device: {self.device}", flush=True)
print("Loading PAPILA data...", flush=True)
self.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=getattr(args, "iop_corr_method", "ratio"),
iop_drop_raw=getattr(args, "iop_drop_raw", False),
exclude_cols=list(getattr(args, "exclude_cols", []) or []),
)
print(f"Loaded: {len(self.data.df)} rows feature_dim={self.data.feature_dim}", flush=True)
self.image_preprocessor = build_image_preprocessor_from_args(args)
self.profile_eye = build_papila_profile(
patient_col="Patient ID", label_col=args.label_col, sample_mode="eye"
)
self.profile_patient = build_papila_profile(
patient_col="Patient ID", label_col=args.label_col, sample_mode="patient"
)
def run(self) -> Path:
"""Execute the full fold loop."""
args = self.args
ts = time.strftime("%Y%m%d_%H%M%S")
run_name = args.run_name or f"v3_hypertower_{ts}"
out_dir = Path(args.output_root) / run_name
out_dir.mkdir(parents=True, exist_ok=True)
mode = args.eval_mode
tower_mode = "single" if args.tower_mode == "classic" else args.tower_mode
df_mode = self.data.df.copy()
if args.exclude_mixed_patients:
before = df_mode["Patient ID"].nunique()
df_mode, mixed = _drop_mixed_label_patients(
df_mode, patient_col="Patient ID", label_col=args.label_col
)
print(f"[{mode}] dropped {len(mixed)} mixed-label patients ({before}{df_mode['Patient ID'].nunique()})", flush=True)
elif args.relabel_mixed_patients_to_max:
df_mode, changed, still_mixed = _relabel_mixed_patients_to_max(
df_mode, patient_col="Patient ID", label_col=args.label_col
)
print(f"[{mode}] relabeled {changed} mixed-patient rows to max severity", flush=True)
if mode == "binary":
df_mode = df_mode[df_mode[args.label_col].isin([0, 1])].reset_index(drop=True)
num_classes = 2 if mode == "binary" else int(df_mode[args.label_col].nunique())
print(f"\n[{mode}] num_classes={num_classes} rows={len(df_mode)} patients={df_mode['Patient ID'].nunique()}", flush=True)
split_manager = PatientFirstSplitManager(patient_col="Patient ID", label_col=args.label_col)
split_args = SimpleNamespace(
eval_mode=mode,
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)
requested_folds = args.n_splits if args.folds is None else int(args.folds)
n_folds = min(requested_folds, len(plans))
tm_dir = out_dir / mode / tower_mode
tm_dir.mkdir(parents=True, exist_ok=True)
fold_results: list[FoldResult] = []
profile_eye = build_papila_profile(patient_col="Patient ID", label_col=args.label_col, sample_mode="eye")
profile_patient = build_papila_profile(patient_col="Patient ID", label_col=args.label_col, sample_mode="patient")
fused_head = getattr(args, "fused_head", False)
_head_names = head_names_for_mode(tower_mode, fused_head=fused_head)
fusion_epochs = int(getattr(args, "fusion_epochs", 10)) if fused_head else 0
_warmup_tower = int(getattr(args, "single_warmup_tower_epochs", None) or getattr(args, "warmup_tower_epochs", None) or 2)
_warmup_fused = int(getattr(args, "single_warmup_fused_epochs", None) or getattr(args, "warmup_fused_epochs", None) or 2)
_warmup_cd = int(getattr(args, "warmup_cd_epochs", 0))
_total_epochs = _warmup_cd + _warmup_tower + _warmup_fused + int(args.epochs) + fusion_epochs
if tower_mode in ("single", "classic"):
_sample_ids = [f"{row['Patient ID']}{row['eyeID']}" for _, row in df_mode.iterrows()]
_y_true = df_mode[args.label_col].tolist()
else:
_pat_df = df_mode.drop_duplicates(subset="Patient ID")
_sample_ids = _pat_df["Patient ID"].astype(str).tolist()
_y_true = _pat_df[args.label_col].tolist()
pred_store = PredictionStore(
sample_ids=_sample_ids, y_true=_y_true,
head_names=_head_names, n_folds=n_folds,
n_epochs=_total_epochs, n_classes=num_classes,
)
image_cache = CachedImageLoader(
enabled=getattr(args, "in_memory_cache", False),
workers=int(getattr(args, "cache_workers", 4)),
)
for fold in range(n_folds):
seed_everything(args.seed + fold * 100)
fold_dir = tm_dir / f"fold{fold}"
fold_dir.mkdir(exist_ok=True)
print(f"\n[{mode}:{tower_mode}] fold {fold+1}/{n_folds}", flush=True)
result, artifacts = self._run_fold(
fold=fold, split=plans[fold], mode=mode, data=self.data,
num_classes=num_classes, profile_eye=profile_eye,
profile_patient=profile_patient, fold_dir=fold_dir,
tower_mode=tower_mode, pred_store=pred_store, image_cache=image_cache,
)
fold_results.append(result)
# Save val artifacts
if artifacts.y_true_ensemble is not None:
np.save(fold_dir / "y_true.npy", artifacts.y_true_ensemble)
if artifacts.probs_ensemble is not None:
np.save(fold_dir / "probs_fused.npy", artifacts.probs_ensemble)
if artifacts.probs_ensemble_img is not None:
np.save(fold_dir / "probs_img.npy", artifacts.probs_ensemble_img)
if artifacts.probs_ensemble_md is not None:
np.save(fold_dir / "probs_cd.npy", artifacts.probs_ensemble_md)
if artifacts.y_true_classic is not None:
np.save(fold_dir / "y_true.npy", artifacts.y_true_classic)
if artifacts.probs_classic is not None:
np.save(fold_dir / "probs_classic.npy", artifacts.probs_classic)
if artifacts.y_true_ensemble_pereye is not None:
np.save(fold_dir / "y_true_pereye.npy", artifacts.y_true_ensemble_pereye)
if artifacts.probs_ensemble_pereye is not None:
np.save(fold_dir / "probs_fused_pereye.npy", artifacts.probs_ensemble_pereye)
if artifacts.logits_ensemble is not None:
np.save(fold_dir / "logits_fused.npy", artifacts.logits_ensemble)
if artifacts.logits_ensemble_img is not None:
np.save(fold_dir / "logits_img.npy", artifacts.logits_ensemble_img)
if artifacts.logits_ensemble_md is not None:
np.save(fold_dir / "logits_cd.npy", artifacts.logits_ensemble_md)
# Save test artifacts
if artifacts.y_true_test is not None:
np.save(fold_dir / "test_y_true.npy", artifacts.y_true_test)
if artifacts.probs_test is not None:
np.save(fold_dir / "test_probs_fused.npy", artifacts.probs_test)
if artifacts.probs_test_img is not None:
np.save(fold_dir / "test_probs_img.npy", artifacts.probs_test_img)
if artifacts.probs_test_md is not None:
np.save(fold_dir / "test_probs_cd.npy", artifacts.probs_test_md)
if pred_store is not None:
pred_store.save(tm_dir / "predictions.npz")
# Summary
summary = self._summary(fold_results)
self._print_summary(mode, summary, tower_mode=tower_mode)
# Save fold results CSV
import dataclasses
fold_csv = tm_dir / "fold_results.csv"
rows = [dataclasses.asdict(r) for r in fold_results]
if rows:
with fold_csv.open("w", newline="") as f:
writer = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
writer.writeheader()
writer.writerows(rows)
# Save summary JSON
if tower_mode in ("single", "classic"):
_test_key = "classic_test"
elif tower_mode == "ensemble":
_test_key = "ensemble_test"
elif tower_mode == "bilateral":
_test_key = "bilat_test"
else:
_test_key = "classic_test"
with (tm_dir / "summary.json").open("w") as f:
json.dump({"mode_summary": {_test_key: summary.get(_test_key, {})}},
f, indent=2, default=str)
return out_dir
def _run_fold(
self,
*,
fold: int,
split,
mode: str,
data,
num_classes: int,
profile_eye,
profile_patient,
fold_dir: Path,
tower_mode: str,
pred_store,
image_cache,
):
args = self.args
device = self.device
image_preprocessor = self.image_preprocessor
nan = float("nan")
run_single = tower_mode in ("single", "ensemble")
run_bilat = tower_mode == "bilateral"
run_fused = tower_mode == "ensemble" and getattr(args, "fused_head", False)
# ---- warmup schedule -------------------------------------------
global_warmup_tower = getattr(args, "warmup_tower_epochs", None)
global_warmup_fused = getattr(args, "warmup_fused_epochs", None)
single_warmup_tower = int(
getattr(args, "single_warmup_tower_epochs", None) or global_warmup_tower or 2
)
single_warmup_fused = int(
getattr(args, "single_warmup_fused_epochs", None) or global_warmup_fused or 2
)
bilat_warmup_tower = int(
getattr(args, "bilat_warmup_tower_epochs", None) or global_warmup_tower or 3
)
bilat_warmup_fused = int(
getattr(args, "bilat_warmup_fused_epochs", None) or global_warmup_fused or 3
)
single_warmup_cd = int(getattr(args, "warmup_cd_epochs", 0)) if run_single else 0
if not run_single:
single_warmup_tower = single_warmup_fused = 0
if not run_bilat:
bilat_warmup_tower = bilat_warmup_fused = 0
# Warmup is meaningless in single-pathway modes — skip it entirely
_bridge_mode = getattr(args, "bridge_mode", "fused")
if _bridge_mode in ("image_only", "clinical_only"):
single_warmup_cd = single_warmup_tower = single_warmup_fused = 0
bilat_warmup_tower = bilat_warmup_fused = 0
main_epochs = int(args.epochs)
total_single_epochs = (single_warmup_cd + single_warmup_tower + single_warmup_fused + main_epochs) if run_single else 0
total_bilat_epochs = (bilat_warmup_tower + bilat_warmup_fused + main_epochs) if run_bilat else 0
total_epochs = max(total_single_epochs, total_bilat_epochs)
# ---- samples ---------------------------------------------------
eye_train = filter_eye_samples(profile_eye.build_samples(df=split.train, clinical=data))
bilat_train = filter_bilateral_samples(profile_patient.build_samples(df=split.train, clinical=data))
bilat_val = filter_bilateral_samples(profile_patient.build_samples(df=split.val, clinical=data))
bilat_test = filter_bilateral_samples(profile_patient.build_samples(df=split.test, clinical=data)) if split.test is not None else []
if pred_store is not None:
if tower_mode in ("single", "classic"):
train_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in eye_train]
val_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in bilat_val]
else:
train_sids = [str(s["id_1"]) for s in bilat_train]
val_sids = [str(s["id_1"]) for s in bilat_val]
pred_store.set_split(fold, train_sids, "train")
pred_store.set_split(fold, val_sids, "val")
if bilat_test:
test_sids = [str(s["id_1"]) for s in bilat_test]
pred_store.set_split(fold, test_sids, "test")
if len(bilat_val) == 0:
print(f" [fold {fold+1}] WARNING: no bilateral val samples; skipping fold.", flush=True)
empty = FoldResult(
mode=mode, fold=fold,
best_epoch_single=0, best_epoch_bilat=0,
classic_val_auc=nan, classic_val_acc=nan, classic_val_kappa=nan,
classic_val_mcc=nan, classic_val_f1=nan, classic_val_recall=None,
classic_val_ece=nan, classic_val_threshold=nan, classic_val_bias=None,
classic_val_n=0,
ensemble_val_auc=nan, ensemble_val_acc=nan, ensemble_val_kappa=nan,
ensemble_val_mcc=nan, ensemble_val_f1=nan, ensemble_val_recall=None,
ensemble_val_ece=nan, ensemble_val_threshold=nan, ensemble_val_bias=None,
ensemble_val_n=0,
bilat_val_auc=nan, bilat_val_acc=nan, bilat_val_kappa=nan,
bilat_val_mcc=nan, bilat_val_f1=nan, bilat_val_recall=None,
bilat_val_ece=nan, bilat_val_threshold=nan, bilat_val_bias=None,
bilat_val_n=0,
ensemble_test_auc=nan, ensemble_test_acc=nan,
classic_test_auc=nan, classic_test_acc=nan,
bilat_test_auc=nan, bilat_test_acc=nan,
test_n=0,
single_train_n=len(eye_train), bilat_train_n=len(bilat_train),
)
return empty, FoldArtifacts(
y_true_classic=None, probs_classic=None,
y_true_ensemble=None, probs_ensemble=None,
y_true_bilat=None, probs_bilat=None,
)
# ---- models ----------------------------------------------------
single = None
bilateral = None
if run_single:
single = SingleEyeHT(
backbone=args.backbone, freeze_ratio=args.freeze_ratio,
augment=args.augment, clinical_data=data, num_classes=num_classes,
cd_hidden_dim=args.cd_hidden_dim, fusion_dim=args.fusion_dim,
bridge_mode=getattr(args, "bridge_mode", "fused"),
).to(device)
if run_bilat:
bilateral = BilateralHT(
backbone=args.backbone, freeze_ratio=args.freeze_ratio,
augment=args.augment, clinical_data=data, num_classes=num_classes,
cd_hidden_dim=args.cd_hidden_dim, fusion_dim=args.fusion_dim,
).to(device)
slots_eye = profile_eye.slot_descriptors()
slots_patient = profile_patient.slot_descriptors()
loader_kw = dict(batch_size=args.batch_size, num_workers=args.num_workers,
image_cache=image_cache)
# ---- loaders ---------------------------------------------------
use_balanced = bool(getattr(args, "balanced_sampling", False))
train_single_loader = train_eval_loader = train_bilat_loader = cd_only_loader = None
if run_single:
single_sampler = build_balanced_sampler(eye_train) if use_balanced else None
train_single_loader = make_loader(
eye_train, slots_eye, image_transform=single.transform,
image_preprocessor=image_preprocessor, shuffle=True,
sampler=single_sampler, **loader_kw,
)
train_eval_loader = make_loader(
eye_train, slots_eye, image_transform=build_eval_transform(args.backbone),
image_preprocessor=image_preprocessor, shuffle=False, **loader_kw,
)
if single_warmup_cd > 0:
slots_cd_only = {k: v for k, v in slots_eye.items() if k != "image_1"}
md_sampler = single_sampler if single_sampler is not None else build_balanced_sampler(eye_train)
cd_only_loader = make_loader(
eye_train, slots_cd_only, image_transform=None,
image_preprocessor=None, shuffle=True, sampler=md_sampler, **loader_kw,
)
if run_bilat:
bilat_sampler = build_balanced_sampler(bilat_train) if use_balanced else None
train_bilat_loader = make_loader(
bilat_train, slots_patient, image_transform=bilateral.transform,
image_preprocessor=image_preprocessor, shuffle=True,
sampler=bilat_sampler, **loader_kw,
)
elif run_fused:
fused_sampler = build_balanced_sampler(bilat_train) if use_balanced else None
train_bilat_loader = make_loader(
bilat_train, slots_patient, image_transform=single.transform,
image_preprocessor=image_preprocessor, shuffle=True,
sampler=fused_sampler, **loader_kw,
)
eval_transform = build_eval_transform(args.backbone)
val_loader = make_loader(
bilat_val, slots_patient, image_transform=eval_transform,
image_preprocessor=image_preprocessor, shuffle=False, **loader_kw,
)
# ---- test loader (never touched during training) ---------------
test_loader = None
if bilat_test:
test_loader = make_loader(
bilat_test, slots_patient, image_transform=eval_transform,
image_preprocessor=image_preprocessor, shuffle=False, **loader_kw,
)
print(f" [fold {fold+1}] test_n={len(bilat_test)} (bilateral patients)", flush=True)
# ---- prebuild image cache --------------------------------------
for _ldr in [train_single_loader, train_bilat_loader, val_loader, test_loader]:
if _ldr is not None:
_ldr.dataset.prebuild_image_cache()
opt_single = torch.optim.Adam(single.parameters(), lr=args.lr) if run_single else None
opt_bilateral = torch.optim.Adam(bilateral.parameters(), lr=args.lr) if run_bilat else None
# ---- epoch log -------------------------------------------------
epoch_fields = [
"fold", "epoch", "phase_single", "phase_bilat",
"main_epoch_single", "main_epoch_bilat",
"single_active", "bilat_active",
"single_train_loss", "single_train_acc",
"classic_val_auc", "classic_val_acc", "classic_val_n",
"ensemble_val_auc", "ensemble_val_acc", "ensemble_val_n",
"bilat_train_loss", "bilat_train_acc",
"bilat_val_auc", "bilat_val_acc", "bilat_val_n",
"classic_val_auc_img", "classic_val_acc_img",
"classic_val_auc_cd", "classic_val_acc_cd",
"classic_val_fe_corr", "classic_val_fe_err",
"ensemble_val_auc_img", "ensemble_val_acc_img",
"ensemble_val_auc_cd", "ensemble_val_acc_cd",
"ensemble_val_fe_corr", "ensemble_val_fe_err",
"bilat_val_auc_img", "bilat_val_acc_img",
"bilat_val_auc_cd", "bilat_val_acc_cd",
"bilat_val_fe_corr", "bilat_val_fe_err",
"train_auc_fused", "train_acc_fused",
"train_auc_img", "train_acc_img",
"train_auc_cd", "train_acc_cd",
"train_fe_corr", "train_fe_err",
"train_n",
"is_best_single", "is_best_bilat",
]
if num_classes == 2:
_cm_keys = ["tn", "fp", "fn", "tp"]
else:
_cm_keys = [f"cm_{i}_{j}" for i in range(num_classes) for j in range(num_classes)]
for _split in ("classic_val", "ensemble_val", "train"):
for _head in ("fused", "img", "md"):
for _k in _cm_keys:
epoch_fields.append(f"{_split}_{_head}_{_k}")
fold_logger = HypertowerLogger(run_dir=fold_dir)
# Per-epoch accumulators
_epoch_train_pf: list[np.ndarray] = []
_epoch_train_pi: list[np.ndarray] = []
_epoch_train_pm: list[np.ndarray] = []
_epoch_train_ids: list[np.ndarray] = []
_epoch_train_y: list[np.ndarray] = []
_epoch_val_pf_od: list[np.ndarray] = []
_epoch_val_pi_od: list[np.ndarray] = []
_epoch_val_pm_od: list[np.ndarray] = []
_epoch_val_pf_os: list[np.ndarray] = []
_epoch_val_pi_os: list[np.ndarray] = []
_epoch_val_pm_os: list[np.ndarray] = []
_epoch_val_y: list[np.ndarray] = []
_epoch_val_ids: list[np.ndarray] = []
snap_classic: dict = {}
snap_ensemble: dict = {}
snap_bilat: dict = {}
snap_fused: dict = {}
if run_single:
print(
f" [fold {fold+1}] single_train_n={len(eye_train)} val_n={len(bilat_val)} "
f"test_n={len(bilat_test)} "
f"warmup=md{single_warmup_cd}+twr{single_warmup_tower}+fus{single_warmup_fused} total={total_single_epochs}",
flush=True,
)
else:
print(
f" [fold {fold+1}] bilat_train_n={len(bilat_train)} val_n={len(bilat_val)} "
f"test_n={len(bilat_test)} "
f"bilat_warmup={bilat_warmup_tower}+{bilat_warmup_fused} total={total_bilat_epochs}",
flush=True,
)
_prev_phase_single = "inactive"
# ================================================================
# EPOCH LOOP — no test evaluation during training
# ================================================================
for epoch in range(total_epochs):
_epoch_t0 = time.time()
# Phase logic
if not run_single:
phase_single, main_epoch_single, single_active = "inactive", 0, False
elif epoch < single_warmup_cd:
phase_single, main_epoch_single, single_active = "cd_warmup", 0, True
elif epoch < single_warmup_cd + single_warmup_tower:
phase_single, main_epoch_single, single_active = "tower_warmup", 0, True
elif epoch < single_warmup_cd + single_warmup_tower + single_warmup_fused:
phase_single, main_epoch_single, single_active = "fused_warmup", 0, True
elif epoch < total_single_epochs:
phase_single = "main"
main_epoch_single = epoch - single_warmup_cd - single_warmup_tower - single_warmup_fused + 1
single_active = True
else:
phase_single, main_epoch_single, single_active = "done", main_epochs, False
if not run_bilat:
phase_bilat, main_epoch_bilat, bilat_active = "inactive", 0, False
elif epoch < bilat_warmup_tower:
phase_bilat, main_epoch_bilat, bilat_active = "tower_warmup", 0, True
elif epoch < bilat_warmup_tower + bilat_warmup_fused:
phase_bilat, main_epoch_bilat, bilat_active = "fused_warmup", 0, True
elif epoch < total_bilat_epochs:
phase_bilat = "main"
main_epoch_bilat = epoch - bilat_warmup_tower - bilat_warmup_fused + 1
bilat_active = True
else:
phase_bilat, main_epoch_bilat, bilat_active = "done", main_epochs, False
# Training steps
if run_single and single_active:
_active_loader = cd_only_loader if phase_single == "cd_warmup" else train_single_loader
sl_loss, sl_acc = train_single_epoch(
single, _active_loader, opt_single, device,
phase=phase_single, bcd_prob=float(args.bcd_prob),
tower_loss_mode=args.tower_loss_mode,
)
else:
sl_loss, sl_acc = nan, nan
if run_bilat and bilat_active:
bl_loss, bl_acc = train_bilateral_epoch(
bilateral, train_bilat_loader, opt_bilateral, device,
phase=phase_bilat, bcd_prob=float(args.bcd_prob),
tower_loss_mode=args.tower_loss_mode,
)
else:
bl_loss, bl_acc = nan, nan
_skip_val_eval = (phase_single == "cd_warmup")
# Val evaluation
if run_single and tower_mode == "single" and not _skip_val_eval:
y_cl, p_cl, p_cl_img, p_cl_cd = collect_probs_single_components(
single, val_loader, device, aggregate_patient=False
)
cl_acc, cl_auc, cl_n = _score_arrays(y_cl, p_cl, num_classes)
cl_acc_img = float((p_cl_img.argmax(1)==y_cl).mean()) if y_cl.size else nan
cl_acc_cd = float((p_cl_cd.argmax(1) ==y_cl).mean()) if y_cl.size else nan
_, cl_auc_img, _ = _score_arrays(y_cl, p_cl_img, num_classes)
_, cl_auc_cd, _ = _score_arrays(y_cl, p_cl_cd, num_classes)
y_en = np.array([], dtype=np.int64)
p_en = p_en_img = p_en_cd = np.zeros((0, num_classes), dtype=np.float32)
en_acc = en_auc = nan; en_n = 0
en_acc_img = en_acc_cd = en_auc_img = en_auc_cd = nan
elif run_single and tower_mode == "ensemble" and not _skip_val_eval:
(y_en, _p_en_f_od, _p_en_i_od, _p_en_m_od,
_p_en_f_os, _p_en_i_os, _p_en_m_os,
_en_pat_ids) = collect_probs_ensemble_pereye(
single, val_loader, device, return_ids=True
)
p_en = 0.5 * (_p_en_f_od + _p_en_f_os)
p_en_img = 0.5 * (_p_en_i_od + _p_en_i_os)
p_en_cd = 0.5 * (_p_en_m_od + _p_en_m_os)
en_acc, en_auc, en_n = _score_arrays(y_en, p_en, num_classes)
en_acc_img = float((p_en_img.argmax(1)==y_en).mean()) if y_en.size else nan
en_acc_cd = float((p_en_cd.argmax(1) ==y_en).mean()) if y_en.size else nan
_, en_auc_img, _ = _score_arrays(y_en, p_en_img, num_classes)
_, en_auc_cd, _ = _score_arrays(y_en, p_en_cd, num_classes)
y_cl = np.array([], dtype=np.int64)
p_cl = p_cl_img = p_cl_cd = np.zeros((0, num_classes), dtype=np.float32)
cl_acc = cl_auc = nan; cl_n = 0
cl_acc_img = cl_acc_cd = cl_auc_img = cl_auc_cd = nan
else:
y_cl = y_en = np.array([], dtype=np.int64)
p_cl = p_cl_img = p_cl_cd = np.zeros((0, num_classes), dtype=np.float32)
p_en = p_en_img = p_en_cd = np.zeros((0, num_classes), dtype=np.float32)
cl_acc = cl_auc = en_acc = en_auc = nan; cl_n = en_n = 0
cl_acc_img = cl_acc_cd = en_acc_img = en_acc_cd = nan
cl_auc_img = cl_auc_cd = en_auc_img = en_auc_cd = nan
if run_bilat and not _skip_val_eval:
y_bi, p_bi, p_bi_img, p_bi_cd = collect_probs_bilateral_components(
bilateral, val_loader, device
)
bi_acc, bi_auc, bi_n = _score_arrays(y_bi, p_bi, num_classes)
bi_acc_img = float((p_bi_img.argmax(1)==y_bi).mean()) if y_bi.size else nan
bi_acc_cd = float((p_bi_cd.argmax(1) ==y_bi).mean()) if y_bi.size else nan
_, bi_auc_img, _ = _score_arrays(y_bi, p_bi_img, num_classes)
_, bi_auc_cd, _ = _score_arrays(y_bi, p_bi_cd, num_classes)
else:
y_bi = np.array([], dtype=np.int64)
p_bi = np.zeros((0, 0), dtype=np.float32)
bi_acc = bi_auc = nan; bi_n = 0
bi_acc_img = bi_acc_cd = bi_auc_img = bi_auc_cd = nan
# Fusion events
cl_fe_corr, cl_fe_err = _fusion_events(y_cl, p_cl, p_cl_img, p_cl_cd) if y_cl.size else (0, 0)
en_fe_corr, en_fe_err = _fusion_events(y_en, p_en, p_en_img, p_en_cd) if y_en.size else (0, 0)
bi_fe_corr, bi_fe_err = (0, 0)
# Train eval pass
tr_auc_f = tr_acc_f = tr_auc_i = tr_acc_i = tr_auc_m = tr_acc_m = nan
tr_fe_corr = tr_fe_err = tr_n = 0
y_tr = np.array([], dtype=np.int64)
p_tr_f = p_tr_i = p_tr_m = np.zeros((0, num_classes), dtype=np.float32)
if run_single and train_eval_loader is not None and not _skip_val_eval:
y_tr, p_tr_f, p_tr_i, p_tr_m, tr_ids = collect_probs_eye_level(
single, train_eval_loader, device, return_ids=True
)
if y_tr.size:
_, tr_auc_f, _ = _score_arrays(y_tr, p_tr_f, num_classes)
tr_acc_f = float((p_tr_f.argmax(1)==y_tr).mean())
_, tr_auc_i, _ = _score_arrays(y_tr, p_tr_i, num_classes)
tr_acc_i = float((p_tr_i.argmax(1)==y_tr).mean())
_, tr_auc_m, _ = _score_arrays(y_tr, p_tr_m, num_classes)
tr_acc_m = float((p_tr_m.argmax(1)==y_tr).mean())
tr_fe_corr, tr_fe_err = _fusion_events(y_tr, p_tr_f, p_tr_i, p_tr_m)
tr_n = int(y_tr.size)
_epoch_train_pf.append(p_tr_f)
_epoch_train_pi.append(p_tr_i)
_epoch_train_pm.append(p_tr_m)
_epoch_train_ids.append(tr_ids)
_epoch_train_y.append(y_tr)
if pred_store is not None:
if tower_mode in ("single", "classic"):
pred_store.record(fold, epoch, tr_ids, "fused", p_tr_f)
pred_store.record(fold, epoch, tr_ids, "img", p_tr_i)
pred_store.record(fold, epoch, tr_ids, "md", p_tr_m)
else:
od_mask = np.array([str(i).endswith("OD") for i in tr_ids])
os_mask = ~od_mask
od_pids = [str(i)[:-2] for i in tr_ids[od_mask]]
os_pids = [str(i)[:-2] for i in tr_ids[os_mask]]
pred_store.record(fold, epoch, od_pids, "od_fused", p_tr_f[od_mask])
pred_store.record(fold, epoch, od_pids, "od_img", p_tr_i[od_mask])
pred_store.record(fold, epoch, od_pids, "od_md", p_tr_m[od_mask])
pred_store.record(fold, epoch, os_pids, "os_fused", p_tr_f[os_mask])
pred_store.record(fold, epoch, os_pids, "os_img", p_tr_i[os_mask])
pred_store.record(fold, epoch, os_pids, "os_md", p_tr_m[os_mask])
# Accumulate val per-epoch npy
if run_single and tower_mode == "ensemble" and y_en.size:
_epoch_val_pf_od.append(_p_en_f_od); _epoch_val_pi_od.append(_p_en_i_od)
_epoch_val_pm_od.append(_p_en_m_od); _epoch_val_pf_os.append(_p_en_f_os)
_epoch_val_pi_os.append(_p_en_i_os); _epoch_val_pm_os.append(_p_en_m_os)
_epoch_val_y.append(y_en); _epoch_val_ids.append(_en_pat_ids)
elif run_single and tower_mode == "single" and y_cl.size:
_epoch_val_pf_od.append(p_cl); _epoch_val_pi_od.append(p_cl_img)
_epoch_val_pm_od.append(p_cl_cd); _epoch_val_pf_os.append(p_cl)
_epoch_val_pi_os.append(p_cl_img); _epoch_val_pm_os.append(p_cl_cd)
_epoch_val_y.append(y_cl)
# Val PredictionStore
if pred_store is not None:
if run_single and tower_mode == "ensemble" and y_en.size:
pred_store.record(fold, epoch, _en_pat_ids, "od_fused", _p_en_f_od)
pred_store.record(fold, epoch, _en_pat_ids, "od_img", _p_en_i_od)
pred_store.record(fold, epoch, _en_pat_ids, "od_md", _p_en_m_od)
pred_store.record(fold, epoch, _en_pat_ids, "os_fused", _p_en_f_os)
pred_store.record(fold, epoch, _en_pat_ids, "os_img", _p_en_i_os)
pred_store.record(fold, epoch, _en_pat_ids, "os_md", _p_en_m_os)
is_best_single = False
is_best_bilat = False
# Per-epoch test eval (logged only, never used for model selection)
te_auc = te_acc = float("nan")
if test_loader is not None and phase_single == "main":
if run_single and tower_mode == "single":
_yte, _pte, _, _ = collect_probs_single_components(
single, test_loader, device, aggregate_patient=False)
elif run_single and tower_mode == "ensemble":
_yte, _pte, _, _ = collect_probs_single_components(
single, test_loader, device, aggregate_patient=True)
elif run_bilat:
_yte, _pte, _, _ = collect_probs_bilateral_components(
bilateral, test_loader, device)
else:
_yte = _pte = None
if _yte is not None and _yte.size and len(np.unique(_yte)) > 1:
te_auc = float(_score_arrays(_yte, _pte, num_classes)[1])
te_acc = float((_pte.argmax(1) == _yte).mean())
# Confusion matrix cells
def _prefixed_cm(prefix, y, pf, pi, pm):
out = {}
for head, p in (("fused", pf), ("img", pi), ("md", pm)):
for k, v in _cm_cells(y, p, num_classes).items():
out[f"{prefix}_{head}_{k}"] = v
return out
cm_row = {}
cm_row.update(_prefixed_cm("classic_val", y_cl, p_cl, p_cl_img, p_cl_cd))
cm_row.update(_prefixed_cm("ensemble_val", y_en, p_en, p_en_img, p_en_cd))
cm_row.update(_prefixed_cm("train", y_tr, p_tr_f, p_tr_i, p_tr_m))
fold_logger.write_epoch_row({
"fold": fold, "epoch": epoch + 1,
"phase_single": phase_single, "phase_bilat": phase_bilat,
"main_epoch_single": main_epoch_single, "main_epoch_bilat": main_epoch_bilat,
"single_active": int(single_active), "bilat_active": int(bilat_active),
"single_train_loss": _f(sl_loss), "single_train_acc": _f(sl_acc),
"classic_val_auc": _f(cl_auc), "classic_val_acc": _f(cl_acc), "classic_val_n": cl_n,
"ensemble_val_auc": _f(en_auc), "ensemble_val_acc": _f(en_acc), "ensemble_val_n": en_n,
"bilat_train_loss": _f(bl_loss), "bilat_train_acc": _f(bl_acc),
"bilat_val_auc": _f(bi_auc), "bilat_val_acc": _f(bi_acc), "bilat_val_n": bi_n,
"classic_val_auc_img": _f(cl_auc_img), "classic_val_acc_img": _f(cl_acc_img),
"classic_val_auc_cd": _f(cl_auc_cd), "classic_val_acc_cd": _f(cl_acc_cd),
"classic_val_fe_corr": cl_fe_corr, "classic_val_fe_err": cl_fe_err,
"ensemble_val_auc_img": _f(en_auc_img), "ensemble_val_acc_img": _f(en_acc_img),
"ensemble_val_auc_cd": _f(en_auc_cd), "ensemble_val_acc_cd": _f(en_acc_cd),
"ensemble_val_fe_corr": en_fe_corr, "ensemble_val_fe_err": en_fe_err,
"bilat_val_auc_img": _f(bi_auc_img), "bilat_val_acc_img": _f(bi_acc_img),
"bilat_val_auc_cd": _f(bi_auc_cd), "bilat_val_acc_cd": _f(bi_acc_cd),
"bilat_val_fe_corr": bi_fe_corr, "bilat_val_fe_err": bi_fe_err,
"train_auc_fused": _f(tr_auc_f), "train_acc_fused": _f(tr_acc_f),
"train_auc_img": _f(tr_auc_i), "train_acc_img": _f(tr_acc_i),
"train_auc_cd": _f(tr_auc_m), "train_acc_cd": _f(tr_acc_m),
"train_fe_corr": tr_fe_corr, "train_fe_err": tr_fe_err,
"train_n": tr_n,
"is_best_single": int(is_best_single),
"is_best_bilat": int(is_best_bilat),
"test_auc": _f(te_auc), "test_acc": _f(te_acc),
**cm_row,
}, optional_cols=epoch_fields)
# md_warmup progress bar
if phase_single == "cd_warmup":
_bar_w = 30
_filled = int(_bar_w * (epoch + 1) / single_warmup_cd)
_bar = "#" * _filled + "-" * (_bar_w - _filled)
msg = f" [fold {fold+1}] md_warmup [{_bar}] {epoch+1}/{single_warmup_cd} loss={sl_loss:.4f}"
print(f"\r{msg}", end="", flush=True)
fold_logger.info(msg)
_prev_phase_single = phase_single
continue
if _prev_phase_single == "cd_warmup":
print()
if args.log_every > 0 and (epoch + 1) % args.log_every == 0:
_epoch_secs = time.time() - _epoch_t0
if tower_mode == "ensemble":
msg = f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={en_auc:.4f} acc={en_acc:.4f}"
elif tower_mode == "bilateral":
msg = f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={bi_auc:.4f} acc={bi_acc:.4f}"
else:
msg = f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={cl_auc:.4f} acc={cl_acc:.4f}"
print(msg, flush=True)
fold_logger.info(msg)
_prev_phase_single = phase_single
fold_logger.close()
# No checkpoint saving in V3.
# Save per-epoch npy tensors
if _epoch_train_pf:
ids_ref = _epoch_train_ids[0]; y_ref = _epoch_train_y[0]
np.save(fold_dir / "train_patient_ids.npy", ids_ref)
np.save(fold_dir / "train_y_true.npy", y_ref)
np.save(fold_dir / "train_probs_fused.npy", np.stack(_epoch_train_pf))
np.save(fold_dir / "train_probs_img.npy", np.stack(_epoch_train_pi))
np.save(fold_dir / "train_probs_cd.npy", np.stack(_epoch_train_pm))
if _epoch_val_pf_od:
np.save(fold_dir / "val_y_true_epochs.npy", np.stack(_epoch_val_y))
np.save(fold_dir / "val_probs_fused_od_epochs.npy", np.stack(_epoch_val_pf_od))
np.save(fold_dir / "val_probs_img_od_epochs.npy", np.stack(_epoch_val_pi_od))
np.save(fold_dir / "val_probs_cd_od_epochs.npy", np.stack(_epoch_val_pm_od))
np.save(fold_dir / "val_probs_fused_os_epochs.npy", np.stack(_epoch_val_pf_os))
np.save(fold_dir / "val_probs_img_os_epochs.npy", np.stack(_epoch_val_pi_os))
np.save(fold_dir / "val_probs_cd_os_epochs.npy", np.stack(_epoch_val_pm_os))
if _epoch_val_ids:
np.save(fold_dir / "val_patient_ids.npy", _epoch_val_ids[0])
# ================================================================
# Phase 2: fused head (ensemble only)
# ================================================================
snap_holdout_fused: dict = {}
if run_fused and single is not None:
for p in single.parameters():
p.requires_grad_(False)
fused = FusedEnsembleHT(single, num_classes).to(device)
opt_fused = torch.optim.Adam(fused.eye_scorer.parameters(), lr=args.lr)
fusion_epochs = int(getattr(args, "fusion_epochs", 10))
print(
f" [fold {fold+1}] Phase 2: fusion head bilat_train_n={len(bilat_train)} epochs={fusion_epochs}",
flush=True,
)
_val_pids_for_store = [str(s["id_1"]) for s in bilat_val]
for fep in range(fusion_epochs):
fu_loss, fu_acc = train_fusion_epoch(fused, train_bilat_loader, opt_fused, device)
y_fu, p_fu = collect_probs_fused(fused, val_loader, device)
fu_auc = _score_arrays(y_fu, p_fu, num_classes)[1]
if pred_store is not None and y_fu.size:
pred_store.record(fold, total_single_epochs + fep, _val_pids_for_store, "bilat_fused", p_fu)
if y_fu.size and not np.isnan(fu_auc):
fu_acc_val = _score_arrays(y_fu, p_fu, num_classes)[0]
snap_fused, _, _, _ = _tune_and_snap(y_fu, p_fu, fu_acc_val, num_classes, args, args.ece_bins)
if (fep + 1) % max(1, args.log_every) == 0:
print(
f" [fold {fold+1}] fusion ep{fep+1:>3} loss={fu_loss:.4f} "
f"val_auc={fu_auc:.4f}",
flush=True,
)
# No checkpoint saving for fused head either.
# ================================================================
# Post-training: evaluate ONCE on val (final-epoch state) then test set
# ================================================================
# Val artifacts
y_en_pe_best = p_en_pe_best = p_en_pe_best_img = p_en_pe_best_md = None
l_en_best = l_en_best_img = l_en_best_md = None
l_cl_best = l_cl_best_img = l_cl_best_md = None
l_en_pe_best = l_en_pe_best_img = l_en_pe_best_md = None
if run_single and tower_mode == "single":
y_cl_best, p_cl_best, p_cl_best_img, p_cl_best_md, \
l_cl_best, l_cl_best_img, l_cl_best_md = collect_probs_single_components(
single, val_loader, device, aggregate_patient=False, return_logits=True
)
y_en_best = p_en_best = p_en_best_img = p_en_best_md = None
elif run_single and tower_mode == "ensemble":
y_en_best, p_en_best, p_en_best_img, p_en_best_md, \
l_en_best, l_en_best_img, l_en_best_md = collect_probs_single_components(
single, val_loader, device, aggregate_patient=True, return_logits=True
)
y_en_pe_best, p_en_pe_best, p_en_pe_best_img, p_en_pe_best_md, \
l_en_pe_best, l_en_pe_best_img, l_en_pe_best_md = collect_probs_single_components(
single, val_loader, device, aggregate_patient=False, return_logits=True
)
y_cl_best = p_cl_best = p_cl_best_img = p_cl_best_md = None
else:
y_cl_best = y_en_best = None
p_cl_best = p_en_best = p_en_best_img = p_en_best_md = None
p_cl_best_img = p_cl_best_md = None
if run_bilat:
y_bi_best, p_bi_best = collect_probs_bilateral(bilateral, val_loader, device)
else:
y_bi_best = p_bi_best = None
# Compute val snaps from final-epoch model state
if run_single and tower_mode == "single" and y_cl_best is not None:
snap_cl, _, _, _ = _tune_and_snap(y_cl_best, p_cl_best, float((p_cl_best.argmax(1) == y_cl_best).mean()), num_classes, args, args.ece_bins)
snap_classic = snap_cl
elif run_single and tower_mode == "ensemble" and y_en_best is not None:
snap_en, _, _, _ = _tune_and_snap(y_en_best, p_en_best, float((p_en_best.argmax(1) == y_en_best).mean()), num_classes, args, args.ece_bins)
snap_ensemble = snap_en
if run_bilat and y_bi_best is not None:
snap_bi, _, _, _ = _tune_and_snap(y_bi_best, p_bi_best, float((p_bi_best.argmax(1) == y_bi_best).mean()), num_classes, args, args.ece_bins)
snap_bilat = snap_bi
y_fu_best = p_fu_best = None
if run_fused and single is not None:
y_fu_best, p_fu_best = collect_probs_fused(fused, val_loader, device)
# Test set evaluation (once, never seen during training)
snap_test: dict = {}
y_test_out = p_test_out = p_test_img_out = p_test_md_out = None
if test_loader is not None:
if run_single and tower_mode == "ensemble":
y_test_out, p_test_out, p_test_img_out, p_test_md_out = collect_probs_single_components(
single, test_loader, device, aggregate_patient=True
)
elif run_single and tower_mode == "single":
y_test_out, p_test_out, p_test_img_out, p_test_md_out = collect_probs_single_components(
single, test_loader, device, aggregate_patient=False
)
elif run_bilat:
y_test_out, p_test_out, _, _ = collect_probs_bilateral_components(
bilateral, test_loader, device
)
if y_test_out is not None and y_test_out.size:
test_acc_raw = float((p_test_out.argmax(1) == y_test_out).mean())
snap_test, _, _, _ = _tune_and_snap(
y_test_out, p_test_out, test_acc_raw, num_classes, args, args.ece_bins
)
print(
f" [fold {fold+1}] TEST "
f"auc={snap_test.get('auc', nan):.4f} "
f"acc={snap_test.get('acc', nan):.4f} "
f"kappa={snap_test.get('kappa', nan):.4f} "
f"f1={snap_test.get('macro_f1', nan):.4f} "
f"ece={snap_test.get('ece', nan):.4f} "
f"n={snap_test.get('n', 0)}",
flush=True,
)
_save_predictions_csv(
fold_dir, mode, y_test_out,
{"fused": p_test_out, "img": p_test_img_out, "md": p_test_md_out},
suffix="_test",
)
else:
print(f" [fold {fold+1}] WARNING: no test samples for this fold.", flush=True)
test_n = snap_test.get("n", 0)
return FoldResult(
mode=mode, fold=fold,
best_epoch_single=0, best_epoch_bilat=0,
classic_val_auc=snap_classic.get("auc", nan),
classic_val_acc=snap_classic.get("acc", nan),
classic_val_kappa=snap_classic.get("kappa", nan),
classic_val_mcc=snap_classic.get("mcc", nan),
classic_val_f1=snap_classic.get("macro_f1", nan),
classic_val_recall=_sv(snap_classic.get("per_class_recall")),
classic_val_ece=snap_classic.get("ece", nan),
classic_val_threshold=snap_classic.get("threshold", nan),
classic_val_bias=_svf(snap_classic.get("bias")),
classic_val_n=snap_classic.get("n", 0),
ensemble_val_auc=snap_ensemble.get("auc", nan),
ensemble_val_acc=snap_ensemble.get("acc", nan),
ensemble_val_kappa=snap_ensemble.get("kappa", nan),
ensemble_val_mcc=snap_ensemble.get("mcc", nan),
ensemble_val_f1=snap_ensemble.get("macro_f1", nan),
ensemble_val_recall=_sv(snap_ensemble.get("per_class_recall")),
ensemble_val_ece=snap_ensemble.get("ece", nan),
ensemble_val_threshold=snap_ensemble.get("threshold", nan),
ensemble_val_bias=_svf(snap_ensemble.get("bias")),
ensemble_val_n=snap_ensemble.get("n", 0),
bilat_val_auc=snap_bilat.get("auc", nan),
bilat_val_acc=snap_bilat.get("acc", nan),
bilat_val_kappa=snap_bilat.get("kappa", nan),
bilat_val_mcc=snap_bilat.get("mcc", nan),
bilat_val_f1=snap_bilat.get("macro_f1", nan),
bilat_val_recall=_sv(snap_bilat.get("per_class_recall")),
bilat_val_ece=snap_bilat.get("ece", nan),
bilat_val_threshold=snap_bilat.get("threshold", nan),
bilat_val_bias=_svf(snap_bilat.get("bias")),
bilat_val_n=snap_bilat.get("n", 0),
ensemble_test_auc=snap_test.get("auc", nan) if tower_mode == "ensemble" else nan,
ensemble_test_acc=snap_test.get("acc", nan) if tower_mode == "ensemble" else nan,
ensemble_test_kappa=snap_test.get("kappa", nan) if tower_mode == "ensemble" else nan,
ensemble_test_f1=snap_test.get("macro_f1", nan) if tower_mode == "ensemble" else nan,
ensemble_test_ece=snap_test.get("ece", nan) if tower_mode == "ensemble" else nan,
classic_test_auc=snap_test.get("auc", nan) if tower_mode == "single" else nan,
classic_test_acc=snap_test.get("acc", nan) if tower_mode == "single" else nan,
classic_test_kappa=snap_test.get("kappa", nan) if tower_mode == "single" else nan,
classic_test_f1=snap_test.get("macro_f1", nan) if tower_mode == "single" else nan,
classic_test_ece=snap_test.get("ece", nan) if tower_mode == "single" else nan,
bilat_test_auc=snap_test.get("auc", nan) if tower_mode == "bilateral" else nan,
bilat_test_acc=snap_test.get("acc", nan) if tower_mode == "bilateral" else nan,
bilat_test_kappa=snap_test.get("kappa", nan) if tower_mode == "bilateral" else nan,
bilat_test_f1=snap_test.get("macro_f1", nan) if tower_mode == "bilateral" else nan,
bilat_test_ece=snap_test.get("ece", nan) if tower_mode == "bilateral" else nan,
test_n=test_n,
single_train_n=len(eye_train),
bilat_train_n=len(bilat_train),
fused_val_auc=snap_fused.get("auc", nan),
fused_val_acc=snap_fused.get("acc", nan),
fused_val_kappa=snap_fused.get("kappa", nan),
fused_val_mcc=snap_fused.get("mcc", nan),
fused_val_f1=snap_fused.get("macro_f1", nan),
fused_val_recall=_sv(snap_fused.get("per_class_recall")),
fused_val_ece=snap_fused.get("ece", nan),
fused_val_threshold=snap_fused.get("threshold", nan),
fused_val_bias=_svf(snap_fused.get("bias")),
fused_val_n=snap_fused.get("n", 0),
fused_test_auc=nan, fused_test_acc=nan,
), FoldArtifacts(
y_true_classic=y_cl_best, probs_classic=p_cl_best,
y_true_ensemble=y_en_best, probs_ensemble=p_en_best,
y_true_bilat=y_bi_best, probs_bilat=p_bi_best,
y_true_fused=y_fu_best, probs_fused=p_fu_best,
probs_ensemble_img=p_en_best_img,
probs_ensemble_md=p_en_best_md,
probs_classic_img=p_cl_best_img,
probs_classic_md=p_cl_best_md,
y_true_ensemble_pereye=y_en_pe_best,
probs_ensemble_pereye=p_en_pe_best,
probs_ensemble_img_pereye=p_en_pe_best_img,
probs_ensemble_md_pereye=p_en_pe_best_md,
logits_ensemble=l_en_best,
logits_ensemble_img=l_en_best_img,
logits_ensemble_md=l_en_best_md,
logits_classic=l_cl_best,
logits_classic_img=l_cl_best_img,
logits_classic_md=l_cl_best_md,
logits_ensemble_pereye=l_en_pe_best,
logits_ensemble_img_pereye=l_en_pe_best_img,
logits_ensemble_md_pereye=l_en_pe_best_md,
y_true_test=y_test_out,
probs_test=p_test_out,
probs_test_img=p_test_img_out,
probs_test_md=p_test_md_out,
)
@staticmethod
def _summary(results: list[FoldResult]) -> dict:
def _ms(vals):
v = np.array([x for x in vals if x is not None and not np.isnan(float(x))], dtype=float)
return (float(np.mean(v)) if v.size else None, float(np.std(v)) if v.size else None)
out = {}
for label, prefix in [
("classic_best_val", "classic_val"),
("ensemble_best_val", "ensemble_val"),
("bilat_best_val", "bilat_val"),
("fused_best_val", "fused_val"),
]:
sub = {}
for m in ["auc", "acc", "kappa", "mcc", "f1", "ece", "threshold"]:
vals = [getattr(r, f"{prefix}_{m}") for r in results]
mean, std = _ms(vals)
sub[f"{m}_mean"] = mean
if m in ("auc", "f1", "kappa"):
sub[f"{m}_std"] = std
out[label] = sub
for label, prefix in [
("ensemble_test", "ensemble_test"),
("classic_test", "classic_test"),
("bilat_test", "bilat_test"),
]:
sub = {}
for m in ["auc", "acc", "kappa", "f1", "ece"]:
vals = [getattr(r, f"{prefix}_{m}") for r in results]
mean, std = _ms(vals)
sub[f"{m}_mean"] = mean
sub[f"{m}_std"] = std
out[label] = sub
for delta_label, prefix_a, prefix_b in [
("delta_ensemble_vs_classic", "classic_val", "ensemble_val"),
("delta_bilat_vs_ensemble", "ensemble_val", "bilat_val"),
]:
delta = {}
for m in ["auc", "f1", "kappa"]:
pairs = [
getattr(r, f"{prefix_b}_{m}") - getattr(r, f"{prefix_a}_{m}")
for r in results
if not np.isnan(float(getattr(r, f"{prefix_a}_{m}")))
and not np.isnan(float(getattr(r, f"{prefix_b}_{m}")))
]
delta[f"{m}_mean"] = float(np.mean(pairs)) if pairs else None
delta[f"{m}_std"] = float(np.std(pairs)) if pairs else None
out[delta_label] = delta
out["n_folds_completed"] = len(results)
return out
@staticmethod
def _print_summary(mode: str, s: dict, tower_mode: str | None = None) -> None:
def f(v):
return " nan " if v is None else f"{v:.4f}"
def fsd(mean, std):
if mean is None: return " nan "
if std is None: return f"{mean:.4f} "
return f"{mean:.4f}±{std:.4f}"
# Resolve test key
if tower_mode in ("single", "classic"):
test_key = "classic_test"
elif tower_mode == "ensemble":
test_key = "ensemble_test"
elif tower_mode == "bilateral":
test_key = "bilat_test"
else:
test_key = "classic_test"
td = s.get(test_key, {})
print(f"\n=== Summary [{mode}] ===")
print(f" {'':26s} {'AUC':>16} {'ACC':>16} {'Kappa':>16} {'F1-mac':>16} {'ECE':>16}")
print(f" {'Test':26s} "
f"{fsd(td.get('auc_mean'), td.get('auc_std')):>16} "
f"{fsd(td.get('acc_mean'), td.get('acc_std')):>16} "
f"{fsd(td.get('kappa_mean'), td.get('kappa_std')):>16} "
f"{fsd(td.get('f1_mean'), td.get('f1_std')):>16} "
f"{fsd(td.get('ece_mean'), td.get('ece_std')):>16}")
print()