1423 lines
73 KiB
Python
1423 lines
73 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.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,
|
|
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(
|
|
"--tower-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.")
|
|
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.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)
|
|
|
|
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 _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
|
|
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()
|