reworked iop_corr, added explainability tools and plotting tools, cleanup codebase

This commit is contained in:
rpotter6298
2026-03-12 16:12:59 +01:00
parent 13ad32683f
commit 6d0698d0fb
35 changed files with 7765 additions and 5723 deletions
+60 -14
View File
@@ -206,6 +206,9 @@ class V2HyperTower:
help="Single-eye model tower warmup (overrides --warmup-tower-epochs).")
ap.add_argument("--single-warmup-fused-epochs", type=int, default=None,
help="Single-eye model fused warmup (overrides --warmup-fused-epochs).")
ap.add_argument("--warmup-md-epochs", type=int, default=0,
help="MD-only warmup epochs before tower warmup. Trains only md_tower + "
"classifier_md (no CNN forward pass, so 50-100 epochs is cheap).")
ap.add_argument("--bilat-warmup-tower-epochs", type=int, default=None,
help="Bilateral model tower warmup (overrides --warmup-tower-epochs).")
ap.add_argument("--bilat-warmup-fused-epochs", type=int, default=None,
@@ -424,7 +427,8 @@ class V2HyperTower:
if getattr(args, "single_warmup_fused_epochs", None) is not None
else int(_global_warmup_fused) if _global_warmup_fused is not None else 2
)
_total_epochs = _warmup_tower + _warmup_fused + int(args.epochs) + fusion_epochs
_warmup_md = int(getattr(args, "warmup_md_epochs", 0))
_total_epochs = _warmup_md + _warmup_tower + _warmup_fused + int(args.epochs) + fusion_epochs
# sample IDs depend on mode: single uses eye IDs, others use patient IDs
if tower_mode in ("single", "classic"):
@@ -572,6 +576,7 @@ class V2HyperTower:
"run_id": run_name,
"backbone": args.backbone,
"epochs": args.epochs,
"warmup_md_epochs": getattr(args, "warmup_md_epochs", 0),
"warmup_tower_epochs": args.warmup_tower_epochs,
"warmup_fused_epochs": args.warmup_fused_epochs,
"single_warmup_tower_epochs": args.single_warmup_tower_epochs,
@@ -659,6 +664,7 @@ class V2HyperTower:
if getattr(args, "bilat_warmup_fused_epochs", None) is not None
else int(global_warmup_fused) if global_warmup_fused is not None else 3
)
single_warmup_md = int(getattr(args, "warmup_md_epochs", 0)) if run_single else 0
if not run_single:
single_warmup_tower = 0
single_warmup_fused = 0
@@ -666,7 +672,7 @@ class V2HyperTower:
bilat_warmup_tower = 0
bilat_warmup_fused = 0
main_epochs = int(args.epochs)
total_single_epochs = (single_warmup_tower + single_warmup_fused + main_epochs) if run_single else 0
total_single_epochs = (single_warmup_md + single_warmup_tower + single_warmup_fused + main_epochs) if run_single else 0
total_bilat_epochs = (bilat_warmup_tower + bilat_warmup_fused + main_epochs) if run_bilat else 0
total_epochs = max(total_single_epochs, total_bilat_epochs)
@@ -744,6 +750,7 @@ class V2HyperTower:
train_single_loader = None
train_eval_loader = None # non-shuffled, no sampler — for per-epoch train logging
train_bilat_loader = None
md_only_loader = None # image-free loader for md_warmup phase
if run_single:
single_sampler = build_balanced_sampler(eye_train) if use_balanced else None
train_single_loader = make_loader(
@@ -761,6 +768,20 @@ class V2HyperTower:
shuffle=False,
**loader_kw,
)
if single_warmup_md > 0:
# MD-only loader: drop image_1 so PIL never opens files during md_warmup.
# Always use balanced sampling for md_warmup — MD features alone are weaker
# than images and collapse to majority class without class balancing.
slots_md_only = {k: v for k, v in slots_eye.items() if k != "image_1"}
md_warmup_sampler = single_sampler if single_sampler is not None else build_balanced_sampler(eye_train)
md_only_loader = make_loader(
eye_train, slots_md_only,
image_transform=None,
image_preprocessor=None,
shuffle=True,
sampler=md_warmup_sampler,
**loader_kw,
)
if run_bilat:
bilat_sampler = build_balanced_sampler(bilat_train) if use_balanced else None
train_bilat_loader = make_loader(
@@ -907,7 +928,7 @@ class V2HyperTower:
print(
f" [fold {fold+1}] single_train_n={len(eye_train)} (eye-level) "
f"val_n={len(bilat_val)} "
f"single_warmup={single_warmup_tower}+{single_warmup_fused} total={total_single_epochs}",
f"single_warmup=md{single_warmup_md}+twr{single_warmup_tower}+fus{single_warmup_fused} total={total_single_epochs}",
flush=True,
)
else:
@@ -919,17 +940,20 @@ class V2HyperTower:
)
# ---- epoch loop ------------------------------------------------
_prev_phase_single = "inactive" # used to detect md_warmup → next phase transition
for epoch in range(total_epochs):
if not run_single:
phase_single, main_epoch_single, single_active = "inactive", 0, False
elif epoch < single_warmup_tower:
elif epoch < single_warmup_md:
phase_single, main_epoch_single, single_active = "md_warmup", 0, True
elif epoch < (single_warmup_md + single_warmup_tower):
phase_single, main_epoch_single, single_active = "tower_warmup", 0, True
elif epoch < (single_warmup_tower + single_warmup_fused):
elif epoch < (single_warmup_md + 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_epoch_single, single_active = (
"main",
epoch - single_warmup_tower - single_warmup_fused + 1,
epoch - single_warmup_md - single_warmup_tower - single_warmup_fused + 1,
True,
)
else:
@@ -951,8 +975,9 @@ class V2HyperTower:
phase_bilat, main_epoch_bilat, bilat_active = "done", main_epochs, False
if run_single and single_active:
_active_loader = md_only_loader if phase_single == "md_warmup" else train_single_loader
sl_loss, sl_acc = train_single_epoch(
single, train_single_loader, opt_single, device,
single, _active_loader, opt_single, device,
phase=phase_single, bcd_prob=float(args.bcd_prob),
)
else:
@@ -966,7 +991,9 @@ class V2HyperTower:
else:
bl_loss, bl_acc = nan, nan
if run_single and tower_mode == "single":
_skip_val_eval = (phase_single == "md_warmup")
if run_single and tower_mode == "single" and not _skip_val_eval:
y_cl, p_cl, p_cl_img, p_cl_md = collect_probs_single_components(
single, val_loader, device, aggregate_patient=False
)
@@ -980,7 +1007,7 @@ class V2HyperTower:
en_acc = en_auc = nan
en_n = 0
en_acc_img = en_acc_md = en_auc_img = en_auc_md = nan
elif run_single and tower_mode == "ensemble":
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,
@@ -1010,7 +1037,7 @@ class V2HyperTower:
cl_acc_img = cl_acc_md = en_acc_img = en_acc_md = nan
cl_auc_img = cl_auc_md = en_auc_img = en_auc_md = nan
if run_bilat:
if run_bilat and not _skip_val_eval:
y_bi, p_bi, p_bi_img, p_bi_md = collect_probs_bilateral_components(
bilateral, val_loader, device
)
@@ -1034,7 +1061,7 @@ class V2HyperTower:
p_cl_h = p_cl_h_img = p_cl_h_md = _z2
p_en_h = p_en_h_img = p_en_h_md = _z2
if holdout_loader is not None:
if holdout_loader is not None and not _skip_val_eval:
if run_single and tower_mode == "single":
y_cl_h, p_cl_h, p_cl_h_img, p_cl_h_md = collect_probs_single_components(
single, holdout_loader, device, aggregate_patient=False
@@ -1092,7 +1119,7 @@ class V2HyperTower:
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:
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
)
@@ -1284,16 +1311,33 @@ class V2HyperTower:
**cm_row,
}, optional_cols=epoch_fields)
# ---- md_warmup progress bar (replaces per-epoch print) --------
if phase_single == "md_warmup":
_bar_w = 30
_filled = int(_bar_w * (epoch + 1) / single_warmup_md)
_bar = "#" * _filled + "-" * (_bar_w - _filled)
_bar_msg = (
f" [fold {fold+1}] md_warmup [{_bar}] "
f"{epoch + 1}/{single_warmup_md} loss={sl_loss:.4f}"
)
print(f"\r{_bar_msg}", end="", flush=True)
fold_logger.info(_bar_msg)
_prev_phase_single = phase_single
continue # skip normal log block entirely
if _prev_phase_single == "md_warmup":
print() # seal the progress bar line
if args.log_every > 0 and (epoch + 1) % args.log_every == 0:
hld_auc = target_holdout_single_auc if run_single else bi_auc_h
hld_suffix = f" hld_auc={hld_auc:.4f}" if holdout_loader is not None else ""
# Human-readable phase progress for console logs.
if phase_single == "tower_warmup":
single_phase_epoch = epoch + 1
single_phase_epoch = epoch - single_warmup_md + 1
single_phase_total = single_warmup_tower
elif phase_single == "fused_warmup":
single_phase_epoch = epoch - single_warmup_tower + 1
single_phase_epoch = epoch - single_warmup_md - single_warmup_tower + 1
single_phase_total = single_warmup_fused
else:
single_phase_epoch = main_epoch_single
@@ -1343,6 +1387,8 @@ class V2HyperTower:
print(msg, flush=True)
fold_logger.info(msg)
_prev_phase_single = phase_single
fold_logger.close()
if args.save_checkpoints: