Files
hypertower/v3/classes/v3_hypertower.py
T
rpotter6298 4dea45df78 v4 update
2026-04-20 18:01:31 +02:00

1745 lines
88 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 (
ManifestImageCropper,
UNetImageCropper,
build_image_preprocessor_from_args,
)
from v3.classes.image_loader import CachedImageLoader
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.bridges import Bridge
from v3.classes.towerbase import train_towers_epoch, collect_probs_towers
from v3.classes.image_towers import ImageTower
from v3.classes.clinical_towers import ClinicalDataTower
from v3.classes.geometry_towers import GeometryTower
from v3.classes.hypertower_models import (
BilateralHT,
EmbeddingMLPEnsembleHT,
FusedEnsembleHT,
LogitMLPEnsembleHT,
SiameseHT,
SingleEyeHT,
V2ModeComparisonOps,
collect_probs_bilateral,
collect_probs_siamese,
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_siamese_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 EyeLevelSplitManager, 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(
"--hypertower-mode", choices=["single", "ensemble", "bilateral", "siamese", "classic"],
default="ensemble",
)
ap.add_argument("--n-splits", type=int, default=5)
ap.add_argument("--fold-seed", type=int, default=42)
ap.add_argument("--leaky-cv", action="store_true",
help="Split at eye level (leaky: same patient can span folds). "
"Used to demonstrate data-leakage effect.")
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=3)
ap.add_argument("--single-warmup-fused-epochs", type=int, default=3)
ap.add_argument("--warmup-cd-epochs", type=int, default=40)
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=8)
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"])
ap.add_argument("--bridge-dropout", type=float, default=0.5,
help="Dropout in bridge classifier_fused (default: 0.5)")
ap.add_argument("--cd-dropout", type=float, default=0.1,
help="Dropout in clinical tower MLP (default: 0.1)")
ap.add_argument("--se-img-tower", action="store_true",
help="Enable SE gate on image tower output features")
ap.add_argument("--se-cd-tower", action="store_true",
help="Enable SE gate on clinical tower output features")
ap.add_argument("--se-bridge", action="store_true",
help="Enable SE gate on fused vector inside the bridge")
# 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)
ap.add_argument("--head-type",
choices=["attention", "logit_mlp", "embedding_mlp"],
default="attention",
help="Which bilateral head to train on top of frozen ensemble base")
ap.add_argument("--save-checkpoints", action="store_true",
help="Save best_single.pt per fold for explainability / GradCAM")
# Geometry features
ap.add_argument("--geometry-dim", type=int, default=0,
help="Append N geometry features to clinical metadata (0=disabled, 5=all). "
"Requires --img-crop-manifest.")
ap.add_argument("--geometry-source", default="gt", choices=["gt", "unet"],
help="Source for geometry features: gt (GT contour annotations) or "
"unet (U-Net segmentation). unet also requires --img-crop-weights.")
ap.add_argument("--geometry-tower", action="store_true",
help="Add a dedicated GeometryTower (disc/cup seg-map CNN) fused via the "
"bridge alongside ImageTower and ClinicalDataTower. Requires "
"--img-crop-manifest.")
ap.add_argument("--geometry-tower-backbone", default="resnet18",
choices=["resnet18", "resnet50", "efficientnet_b0"],
help="SegCNN backbone for GeometryTower (default: resnet18).")
ap.add_argument("--geometry-tower-in-channels", type=int, default=3, choices=[1, 3],
help="1 = single label map; 3 = one-hot disc/rim/cup (default: 3).")
ap.add_argument("--geometry-tower-frozen", action="store_true",
help="Freeze GeometryTower backbone throughout training.")
ap.add_argument("--geometry-tower-finetune-unet-epochs", type=int, default=0,
help="Epochs to fine-tune the U-Net per fold before seg-map extraction "
"(0 = disabled; only applies when --geometry-source unet).")
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"
)
# Build geometry provider if requested, extend feature_dim to include geometry.
# Both ManifestImageCropper and UNetImageCropper already have geometry_features()
# and precompute_geometry() — we just pick the right one and pre-compute upfront.
self.geometry_provider = None
geom_dim = int(getattr(args, "geometry_dim", 0))
if geom_dim > 0:
source = getattr(args, "geometry_source", "gt")
manifest = getattr(args, "img_crop_manifest", None)
if not manifest:
raise ValueError("--geometry-dim requires --img-crop-manifest")
all_paths = [
self.data.get_image_path(row)
for _, row in self.data.df.iterrows()
]
if source == "gt":
# Reuse image_preprocessor if it's already a ManifestImageCropper,
# otherwise build a lightweight one just for geometry (no crop cache).
if isinstance(self.image_preprocessor, ManifestImageCropper):
provider = self.image_preprocessor
else:
provider = ManifestImageCropper(manifest_path=Path(manifest))
print(f"[geometry] GT source — pre-computing geometry from {manifest}", flush=True)
elif source == "unet":
weights = getattr(args, "img_crop_weights", None)
if not weights:
raise ValueError("--geometry-source unet requires --img-crop-weights")
if isinstance(self.image_preprocessor, UNetImageCropper):
provider = self.image_preprocessor
else:
provider = UNetImageCropper(
manifest_path=Path(manifest),
weights_path=Path(weights),
normalize=getattr(args, "img_crop_normalize", "per_image"),
threshold=getattr(args, "img_crop_threshold", 0.5),
)
print(f"[geometry] UNet source — pre-computing geometry from {weights}", flush=True)
else:
raise ValueError(f"Unknown --geometry-source: {source!r}")
provider.precompute_geometry(all_paths)
self.geometry_provider = provider
self.data.feature_dim += geom_dim
print(f"[geometry] feature_dim extended to {self.data.feature_dim} (+{geom_dim} geometry)", flush=True)
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.hypertower_mode == "classic" else args.hypertower_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)
if getattr(args, "leaky_cv", False):
split_manager = EyeLevelSplitManager(patient_col="Patient ID", label_col=args.label_col)
print("[CV] WARNING: leaky-cv mode — eye-level splits, same patient can span folds.", flush=True)
else:
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 in ("bilateral", "siamese"):
_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_towers(
self,
*,
fold: int,
split,
mode: str,
data,
num_classes: int,
profile_eye,
profile_patient,
fold_dir: Path,
pred_store,
image_cache,
):
"""Modular TowerBase training path (used when --geometry-tower is set).
Builds [ImageTower, ClinicalDataTower, GeometryTower], runs the full
fold lifecycle (prepare_fold → augment_samples → loader build → epoch
loop → eval), and returns (FoldResult, FoldArtifacts) with metrics in
the ensemble_val_* slots.
"""
args = self.args
device = self.device
nan = float("nan")
# ------------------------------------------------------------------ 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 []
# Old --geometry-dim path still applies (injects geometry into clinical stream)
if self.geometry_provider is not None:
eye_train = self._augment_geometry(eye_train)
bilat_train = self._augment_geometry(bilat_train)
bilat_val = self._augment_geometry(bilat_val)
bilat_test = self._augment_geometry(bilat_test)
if len(bilat_val) == 0:
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,
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,
)
# ------------------------------------------------------------------ towers
img_tower = ImageTower(
backbone=args.backbone,
freeze_ratio=args.freeze_ratio,
augment=args.augment,
use_se=getattr(args, "se_img_tower", False),
)
cd_tower = ClinicalDataTower(
clinical_data=data,
cd_hidden_dim=args.cd_hidden_dim,
cd_dropout=getattr(args, "cd_dropout", 0.1),
use_se=getattr(args, "se_cd_tower", False),
)
geom_tower = GeometryTower(
backbone=getattr(args, "geometry_tower_backbone", "resnet18"),
in_channels=getattr(args, "geometry_tower_in_channels", 3),
pretrained=not getattr(args, "no_pretrained", False),
frozen=getattr(args, "geometry_tower_frozen", False),
geometry_source=getattr(args, "geometry_source", "gt"),
manifest_path=getattr(args, "img_crop_manifest", None),
weights_path=getattr(args, "img_crop_weights", None),
unet_normalize=getattr(args, "img_crop_normalize", "per_image"),
unet_threshold=getattr(args, "img_crop_threshold", 0.5),
finetune_unet_epochs=getattr(args, "geometry_tower_finetune_unet_epochs", 0),
)
# GeometryTower.prepare_fold must run before augment_samples (precomputes seg maps)
geom_tower.prepare_fold(
eye_train=eye_train, bilat_train=bilat_train,
bilat_val=bilat_val, bilat_test=bilat_test,
image_preprocessor=self.image_preprocessor,
image_cache=image_cache, device=device, args=args,
)
# Inject seg_map_1/seg_map_2 into all sample lists before loaders are built
for sample_list in (eye_train, bilat_train, bilat_val, bilat_test):
geom_tower.augment_samples(sample_list)
# Now ImageTower.prepare_fold sees augmented samples → loader includes seg maps
img_tower.prepare_fold(
eye_train=eye_train, bilat_train=bilat_train,
bilat_val=bilat_val, bilat_test=bilat_test,
image_preprocessor=self.image_preprocessor,
image_cache=image_cache, device=device, args=args,
)
cd_tower.prepare_fold(
eye_train=eye_train, bilat_train=bilat_train,
bilat_val=bilat_val, bilat_test=bilat_test,
image_preprocessor=self.image_preprocessor,
image_cache=image_cache, device=device, args=args,
)
towers = [img_tower, cd_tower, geom_tower]
# ------------------------------------------------------------------ bridge
tower_dims = []
for t in towers:
tower_dims.extend(t.embed_dims)
bridge = Bridge(
tower_dims=tower_dims,
num_classes=num_classes,
fusion_dim=args.fusion_dim,
mode=getattr(args, "bridge_mode", "fused"),
dropout=getattr(args, "bridge_dropout", 0.5),
)
# Move all nn.Modules to device
for t in towers:
if isinstance(t, torch.nn.Module):
t.to(device)
bridge.to(device)
# ------------------------------------------------------------------ loaders
slots_patient = profile_patient.slot_descriptors()
_persistent = args.num_workers > 0
loader_kw = dict(
batch_size=args.batch_size, num_workers=args.num_workers,
image_cache=image_cache, persistent_workers=_persistent,
)
eval_transform = build_eval_transform(args.backbone)
val_loader = make_loader(
bilat_val, slots_patient,
image_transform=eval_transform,
image_preprocessor=self.image_preprocessor,
shuffle=False, **loader_kw,
)
test_loader = None
if bilat_test:
test_loader = make_loader(
bilat_test, slots_patient,
image_transform=eval_transform,
image_preprocessor=self.image_preprocessor,
shuffle=False, **loader_kw,
)
train_loader = img_tower.train_loader
val_loader.dataset.prebuild_image_cache()
if test_loader is not None:
test_loader.dataset.prebuild_image_cache()
# ------------------------------------------------------------------ optimizer
all_params = list(bridge.parameters())
for t in towers:
if isinstance(t, torch.nn.Module):
all_params.extend(t.parameters())
optimizer = torch.optim.AdamW(
[p for p in all_params if p.requires_grad],
lr=args.lr,
weight_decay=getattr(args, "weight_decay", 1e-4),
)
# ------------------------------------------------------------------ epoch loop
global_warmup_tower = getattr(args, "warmup_tower_epochs", None)
global_warmup_fused = getattr(args, "warmup_fused_epochs", None)
warmup_cd = int(getattr(args, "warmup_cd_epochs", 0))
warmup_tower = int(getattr(args, "single_warmup_tower_epochs", None) or global_warmup_tower or 2)
warmup_fused = int(getattr(args, "single_warmup_fused_epochs", None) or global_warmup_fused or 2)
main_epochs = int(args.epochs)
schedule = []
if warmup_cd > 0: schedule.append(("cd_warmup", warmup_cd))
if warmup_tower > 0: schedule.append(("tower_warmup", warmup_tower))
if warmup_fused > 0: schedule.append(("fused_warmup", warmup_fused))
schedule.append(("main", main_epochs))
best_val_auc = float("-inf")
best_epoch = 0
best_tower_states = None
best_bridge_state = None
epoch_idx = 0
for phase, n_epochs in schedule:
for _ in range(n_epochs):
for t in towers:
if isinstance(t, torch.nn.Module):
t.train()
train_towers_epoch(
towers, bridge, train_loader, optimizer, device,
phase=phase,
bcd_prob=getattr(args, "bcd_prob", 0.5),
tower_loss_mode=getattr(args, "tower_loss_mode", "bcd"),
)
y_v, p_v = collect_probs_towers(towers, bridge, val_loader, device,
tower_mode="ensemble")
_, val_auc, _ = _score_arrays(y_v, p_v, num_classes)
if not np.isnan(val_auc) and val_auc > best_val_auc:
best_val_auc = val_auc
best_epoch = epoch_idx
best_tower_states = [
t.state_dict() if isinstance(t, torch.nn.Module) else None
for t in towers
]
best_bridge_state = bridge.state_dict()
epoch_idx += 1
# Restore best
if best_bridge_state is not None:
bridge.load_state_dict(best_bridge_state)
if best_tower_states is not None:
for t, st in zip(towers, best_tower_states):
if isinstance(t, torch.nn.Module) and st is not None:
t.load_state_dict(st)
# ------------------------------------------------------------------ eval
y_val, p_val = collect_probs_towers(towers, bridge, val_loader, device, tower_mode="ensemble")
acc_val, auc_val, n_val = _score_arrays(y_val, p_val, num_classes)
snap_val, _, thr_val, bias_val = _tune_and_snap(
y_val, p_val, acc_val, num_classes, args, n_bins=10
)
y_test = p_test = None
test_auc = test_acc = nan
test_n = 0
if test_loader is not None:
y_test, p_test = collect_probs_towers(towers, bridge, test_loader, device,
tower_mode="ensemble")
test_acc, test_auc, test_n = _score_arrays(y_test, p_test, num_classes)
result = FoldResult(
mode=mode, fold=fold,
best_epoch_single=best_epoch, 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=snap_val["auc"], ensemble_val_acc=snap_val["acc"],
ensemble_val_kappa=snap_val["kappa"], ensemble_val_mcc=snap_val["mcc"],
ensemble_val_f1=snap_val["macro_f1"],
ensemble_val_recall=_sv(snap_val["per_class_recall"]),
ensemble_val_ece=snap_val["ece"],
ensemble_val_threshold=snap_val["threshold"],
ensemble_val_bias=_svf(bias_val),
ensemble_val_n=snap_val["n"],
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=test_auc, ensemble_test_acc=test_acc,
test_n=test_n,
single_train_n=len(eye_train), bilat_train_n=len(bilat_train),
)
artifacts = FoldArtifacts(
y_true_classic=None, probs_classic=None,
y_true_ensemble=y_val, probs_ensemble=p_val,
y_true_bilat=None, probs_bilat=None,
y_true_test=y_test, probs_test=p_test,
)
return result, artifacts
def _augment_geometry_slot(self, samples: list) -> list:
"""Add geom_1/geom_2 keys to each sample dict (geometry tower mode).
Unlike _augment_geometry, this does NOT touch matrix_1/matrix_2 — the
geometry vector lives in its own slot so ImageTower and ClinicalDataTower
each receive only their own modality.
"""
if self.geometry_provider is None:
return samples
geom_dim = int(getattr(self.args, "geometry_dim", 0)) or 5
for s in samples:
for img_slot, geom_slot in (("image_1", "geom_1"), ("image_2", "geom_2")):
img_path = s.get(img_slot)
if img_path is None:
continue
vec = self.geometry_provider.geometry_for_image(img_path)
if vec is not None and len(vec) >= geom_dim:
s[geom_slot] = vec[:geom_dim].astype(np.float32)
else:
s[geom_slot] = np.zeros(geom_dim, dtype=np.float32)
return samples
def _augment_geometry(self, samples: list) -> list:
"""Append geometry features to matrix_1/matrix_2 in each sample dict."""
if self.geometry_provider is None:
return samples
geom_dim = int(getattr(self.args, "geometry_dim", 0))
for s in samples:
for img_slot, mat_slot in (("image_1", "matrix_1"), ("image_2", "matrix_2")):
img_path = s.get(img_slot)
mat = s.get(mat_slot)
if img_path is None or mat is None:
continue
vec = self.geometry_provider.geometry_for_image(img_path)
if vec is not None and len(vec) >= geom_dim:
geom = vec[:geom_dim].astype(np.float32)
else:
geom = np.zeros(geom_dim, dtype=np.float32)
s[mat_slot] = np.concatenate([np.asarray(mat, dtype=np.float32), geom])
return samples
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
# Modular tower path — bypasses the legacy single/bilat/siamese code entirely
if getattr(args, "geometry_tower", False):
return self._run_fold_towers(
fold=fold, split=split, mode=mode, data=data,
num_classes=num_classes, profile_eye=profile_eye,
profile_patient=profile_patient, fold_dir=fold_dir,
pred_store=pred_store, image_cache=image_cache,
)
device = self.device
image_preprocessor = self.image_preprocessor
nan = float("nan")
run_single = tower_mode in ("single", "ensemble")
run_bilat = tower_mode == "bilateral"
run_siamese = tower_mode == "siamese"
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 and not run_siamese:
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 or run_siamese) 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 self.geometry_provider is not None:
eye_train = self._augment_geometry(eye_train)
bilat_train = self._augment_geometry(bilat_train)
bilat_val = self._augment_geometry(bilat_val)
bilat_test = self._augment_geometry(bilat_test)
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 (CPU for now — moved to device after workers spawn) ---
single = None
bilateral = None
siamese = 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"),
bridge_dropout=getattr(args, "bridge_dropout", 0.5),
cd_dropout=getattr(args, "cd_dropout", 0.1),
se_img_tower=getattr(args, "se_img_tower", False),
se_cd_tower=getattr(args, "se_cd_tower", False),
se_bridge=getattr(args, "se_bridge", False),
)
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,
)
if run_siamese:
siamese = SiameseHT(
backbone=args.backbone, freeze_ratio=args.freeze_ratio,
augment=args.augment, num_classes=num_classes,
fusion_dim=args.fusion_dim,
)
slots_eye = profile_eye.slot_descriptors()
slots_patient = profile_patient.slot_descriptors()
_persistent_workers = args.num_workers > 0
loader_kw = dict(batch_size=args.batch_size, num_workers=args.num_workers,
image_cache=image_cache,
persistent_workers=_persistent_workers)
# ---- 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_siamese:
siamese_sampler = build_balanced_sampler(bilat_train) if use_balanced else None
train_bilat_loader = make_loader(
bilat_train, slots_patient, image_transform=siamese.transform,
image_preprocessor=image_preprocessor, shuffle=True,
sampler=siamese_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()
# ---- spawn DataLoader workers BEFORE CUDA init -----------------
# Workers fork here (clean process state, no CUDA context yet).
# persistent_workers=True keeps them alive so the training loop
# reuses them rather than re-forking after .to(device).
if _persistent_workers:
for _ldr in [train_single_loader, train_bilat_loader, val_loader, test_loader]:
if _ldr is not None:
_ = iter(_ldr) # triggers fork now, before CUDA
# ---- move models to device (CUDA init happens here) ------------
if single is not None:
single = single.to(device)
if bilateral is not None:
bilateral = bilateral.to(device)
if siamese is not None:
siamese = siamese.to(device)
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
opt_siamese = torch.optim.Adam(siamese.parameters(), lr=args.lr) if run_siamese 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 and not run_siamese:
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,
)
elif run_siamese and bilat_active:
bl_loss, bl_acc = train_siamese_epoch(
siamese, train_bilat_loader, opt_siamese, device,
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)
elif run_siamese and not _skip_val_eval:
y_bi, p_bi = collect_probs_siamese(siamese, val_loader, device)
bi_acc, bi_auc, bi_n = _score_arrays(y_bi, p_bi, num_classes)
bi_acc_img = bi_acc_cd = bi_auc_img = bi_auc_cd = nan
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:.2f}"
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":
_auc_v, _acc_v = en_auc, en_acc
elif tower_mode in ("bilateral", "siamese"):
_auc_v, _acc_v = bi_auc, bi_acc
else:
_auc_v, _acc_v = cl_auc, cl_acc
print(f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={_auc_v:.2f} acc={_acc_v:.2f}", flush=True)
fold_logger.info(f" ep {epoch+1:>3}/{total_epochs} ({_epoch_secs:.1f}s) auc={_auc_v:.4f} acc={_acc_v:.4f}")
_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)
head_type = getattr(args, "head_type", "attention")
if head_type == "logit_mlp":
fused = LogitMLPEnsembleHT(single, num_classes).to(device)
elif head_type == "embedding_mlp":
fused = EmbeddingMLPEnsembleHT(single, num_classes).to(device)
else:
fused = FusedEnsembleHT(single, num_classes).to(device)
opt_fused = torch.optim.Adam(fused.head.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:.2f} "
f"val_auc={fu_auc:.2f}",
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)
elif run_siamese:
y_bi_best, p_bi_best = collect_probs_siamese(siamese, val_loader, device)
else:
y_bi_best = p_bi_best = None
# Optional checkpoint saving (final-epoch weights for explainability)
if getattr(args, "save_checkpoints", False) and run_single and single is not None:
import torch as _torch
_torch.save(single.state_dict(), fold_dir / "best_single.pt")
# 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 or run_siamese) 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
)
elif run_siamese:
y_test_out, p_test_out = collect_probs_siamese(siamese, 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):.2f} "
f"acc={snap_test.get('acc', nan):.2f} "
f"kappa={snap_test.get('kappa', nan):.2f} "
f"f1={snap_test.get('macro_f1', nan):.2f} "
f"ece={snap_test.get('ece', nan):.2f} "
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 in ("bilateral", "siamese") else nan,
bilat_test_acc=snap_test.get("acc", nan) if tower_mode in ("bilateral", "siamese") else nan,
bilat_test_kappa=snap_test.get("kappa", nan) if tower_mode in ("bilateral", "siamese") else nan,
bilat_test_f1=snap_test.get("macro_f1", nan) if tower_mode in ("bilateral", "siamese") else nan,
bilat_test_ece=snap_test.get("ece", nan) if tower_mode in ("bilateral", "siamese") 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:.2f}"
def fsd(mean, std):
if mean is None: return " nan "
if std is None: return f"{mean:.2f} "
return f"{mean:.2f}±{std:.2f}"
# Resolve test key
if tower_mode in ("single", "classic"):
test_key = "classic_test"
elif tower_mode == "ensemble":
test_key = "ensemble_test"
elif tower_mode in ("bilateral", "siamese"):
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()