"""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()