reworked iop_corr, added explainability tools and plotting tools, cleanup codebase
This commit is contained in:
+60
-14
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user