pre-refactor 041426

This commit is contained in:
rpotter6298
2026-04-14 19:42:16 +02:00
parent eb9eafe715
commit 13290575d5
75 changed files with 8900 additions and 77 deletions
+223 -51
View File
@@ -23,7 +23,11 @@ from typing import Optional
import numpy as np
import torch
from v3.classes.croppers import build_image_preprocessor_from_args
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 (
@@ -35,10 +39,14 @@ from v3.classes.loader_factory import (
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,
@@ -47,6 +55,7 @@ from v3.classes.models import (
collect_probs_fused,
collect_probs_single_components,
train_bilateral_epoch,
train_siamese_epoch,
train_fusion_epoch,
train_single_epoch,
)
@@ -54,7 +63,7 @@ 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.split_manager import EyeLevelSplitManager, PatientFirstSplitManager
from v3.classes.transforms import build_eval_transform
from v3.classes.utils import (
_drop_mixed_label_patients,
@@ -145,11 +154,14 @@ class V3HyperTower:
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"],
"--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.",
@@ -157,9 +169,9 @@ class V3HyperTower:
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("--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)
@@ -170,7 +182,7 @@ class V3HyperTower:
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("--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)
@@ -191,10 +203,20 @@ class V3HyperTower:
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",
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")
@@ -217,6 +239,19 @@ class V3HyperTower:
# 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:
@@ -246,6 +281,49 @@ class V3HyperTower:
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
@@ -276,7 +354,11 @@ class V3HyperTower:
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)
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,
@@ -390,7 +472,7 @@ class V3HyperTower:
_test_key = "classic_test"
elif tower_mode == "ensemble":
_test_key = "ensemble_test"
elif tower_mode == "bilateral":
elif tower_mode in ("bilateral", "siamese"):
_test_key = "bilat_test"
else:
_test_key = "classic_test"
@@ -400,6 +482,25 @@ class V3HyperTower:
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,
*,
@@ -420,9 +521,10 @@ class V3HyperTower:
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)
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)
@@ -442,7 +544,7 @@ class V3HyperTower:
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:
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")
@@ -451,7 +553,7 @@ class V3HyperTower:
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_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 ---------------------------------------------------
@@ -460,6 +562,12 @@ class V3HyperTower:
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]
@@ -502,27 +610,41 @@ class V3HyperTower:
y_true_bilat=None, probs_bilat=None,
)
# ---- models ----------------------------------------------------
# ---- 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"),
).to(device)
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,
).to(device)
)
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)
image_cache=image_cache,
persistent_workers=_persistent_workers)
# ---- loaders ---------------------------------------------------
use_balanced = bool(getattr(args, "balanced_sampling", False))
@@ -553,6 +675,13 @@ class V3HyperTower:
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(
@@ -581,8 +710,26 @@ class V3HyperTower:
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
# ---- 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 = [
@@ -679,7 +826,7 @@ class V3HyperTower:
else:
phase_single, main_epoch_single, single_active = "done", main_epochs, False
if not run_bilat:
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
@@ -709,6 +856,12 @@ class V3HyperTower:
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
@@ -763,6 +916,10 @@ class V3HyperTower:
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)
@@ -905,7 +1062,7 @@ class V3HyperTower:
_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}"
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
@@ -917,13 +1074,13 @@ class V3HyperTower:
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}"
_auc_v, _acc_v = en_auc, en_acc
elif tower_mode in ("bilateral", "siamese"):
_auc_v, _acc_v = bi_auc, bi_acc
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)
_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
@@ -958,8 +1115,14 @@ class V3HyperTower:
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)
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}",
@@ -977,8 +1140,8 @@ class V3HyperTower:
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}",
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.
@@ -1015,9 +1178,16 @@ class V3HyperTower:
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)
@@ -1025,7 +1195,7 @@ class V3HyperTower:
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:
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
@@ -1050,6 +1220,8 @@ class V3HyperTower:
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(
@@ -1057,11 +1229,11 @@ class V3HyperTower:
)
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"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,
)
@@ -1118,11 +1290,11 @@ class V3HyperTower:
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,
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),
@@ -1222,18 +1394,18 @@ class V3HyperTower:
@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}"
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:.4f} "
return f"{mean:.4f}±{std:.4f}"
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 == "bilateral":
elif tower_mode in ("bilateral", "siamese"):
test_key = "bilat_test"
else:
test_key = "classic_test"