pre-refactor 041426
This commit is contained in:
+223
-51
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user