added fused classifier head
This commit is contained in:
+142
-5
@@ -24,6 +24,7 @@ import torch
|
||||
from classes.v2.croppers import build_image_preprocessor_from_args
|
||||
from classes.v2.dataset import _ClinicalView # noqa: F401 (re-exported for compat)
|
||||
from classes.v2.loader_factory import (
|
||||
build_balanced_sampler,
|
||||
filter_bilateral_samples,
|
||||
filter_eye_samples,
|
||||
make_loader,
|
||||
@@ -31,14 +32,17 @@ from classes.v2.loader_factory import (
|
||||
from classes.v2.metrics import _score_arrays, _svf, _tune_and_snap
|
||||
from classes.v2.models import (
|
||||
BilateralHT,
|
||||
FusedEnsembleHT,
|
||||
SingleEyeHT,
|
||||
V2ModeComparisonOps,
|
||||
collect_probs_bilateral,
|
||||
collect_probs_bilateral_components,
|
||||
collect_probs_classic,
|
||||
collect_probs_ensemble,
|
||||
collect_probs_fused,
|
||||
collect_probs_single_components,
|
||||
train_bilateral_epoch,
|
||||
train_fusion_epoch,
|
||||
train_single_epoch,
|
||||
)
|
||||
from classes.v2.papila_builders import build_papila_data
|
||||
@@ -118,6 +122,8 @@ class V2HyperTower:
|
||||
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",
|
||||
help="Use WeightedRandomSampler during training to equalise class frequency (default: off).")
|
||||
ap.add_argument("--num-workers", type=int, default=0)
|
||||
ap.add_argument("--device", choices=["auto", "cpu", "cuda"], default="auto")
|
||||
ap.add_argument("--seed", type=int, default=1234)
|
||||
@@ -186,6 +192,16 @@ class V2HyperTower:
|
||||
ap.add_argument("--log-every", type=int, default=1)
|
||||
ap.add_argument("--save-checkpoints", action=argparse.BooleanOptionalAction, default=True,
|
||||
help="Save best_single.pt / best_holdout_single.pt per fold (use --no-save-checkpoints to disable)")
|
||||
ap.add_argument(
|
||||
"--fused-head", action="store_true",
|
||||
help="(ensemble mode only) After base SingleEyeHT training, freeze it and train a "
|
||||
"small logit-level MLP fusion head on bilateral samples instead of averaging "
|
||||
"OD/OS softmax probabilities.",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--fusion-epochs", type=int, default=10,
|
||||
help="Number of epochs to train the fusion head (--fused-head, ensemble mode only).",
|
||||
)
|
||||
return ap
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -414,7 +430,8 @@ class V2HyperTower:
|
||||
nan = _nan()
|
||||
tower_mode = "single" if tower_mode == "classic" else tower_mode
|
||||
run_single = tower_mode in ("single", "ensemble")
|
||||
run_bilat = tower_mode == "bilateral"
|
||||
run_bilat = tower_mode == "bilateral"
|
||||
run_fused = (tower_mode == "ensemble") and bool(getattr(args, "fused_head", False))
|
||||
|
||||
global_warmup_tower = getattr(args, "warmup_tower_epochs", None)
|
||||
global_warmup_fused = getattr(args, "warmup_fused_epochs", None)
|
||||
@@ -506,22 +523,38 @@ class V2HyperTower:
|
||||
loader_kw = dict(batch_size=args.batch_size, num_workers=args.num_workers)
|
||||
|
||||
# ---- loaders ---------------------------------------------------
|
||||
use_balanced = bool(getattr(args, "balanced_sampling", False))
|
||||
train_single_loader = None
|
||||
train_bilat_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,
|
||||
)
|
||||
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_fused:
|
||||
# Fused head trains on bilateral samples using the single model's transform.
|
||||
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)
|
||||
@@ -579,11 +612,13 @@ class V2HyperTower:
|
||||
best_epoch_bilat = 0
|
||||
best_single_state: Optional[dict] = None
|
||||
best_bilat_state: Optional[dict] = None
|
||||
snap_classic: dict = {}
|
||||
snap_ensemble: dict = {}
|
||||
snap_bilat: dict = {}
|
||||
snap_classic: dict = {}
|
||||
snap_ensemble: dict = {}
|
||||
snap_bilat: dict = {}
|
||||
snap_holdout_single: dict = {}
|
||||
snap_holdout_bilat: dict = {}
|
||||
snap_fused: dict = {}
|
||||
snap_holdout_fused: dict = {}
|
||||
best_holdout_single_auc = -1.0
|
||||
best_holdout_bilat_auc = -1.0
|
||||
best_epoch_holdout_single = 0
|
||||
@@ -863,6 +898,73 @@ class V2HyperTower:
|
||||
if best_holdout_bilat_state is not None:
|
||||
torch.save(best_holdout_bilat_state, fold_dir / "best_holdout_bilateral.pt")
|
||||
|
||||
# ---- Phase 2: fused head training (ensemble + --fused-head only) ----
|
||||
best_fused_auc = -1.0
|
||||
best_fused_state: Optional[dict] = None
|
||||
best_holdout_fused_auc = -1.0
|
||||
|
||||
if run_fused and best_single_state is not None:
|
||||
# Revert base to its best val checkpoint, then freeze it.
|
||||
single.load_state_dict(best_single_state)
|
||||
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)
|
||||
fusion_epochs = int(getattr(args, "fusion_epochs", 10))
|
||||
|
||||
print(
|
||||
f" [fold {fold+1}] Phase 2: training fusion head "
|
||||
f"bilat_train_n={len(bilat_train)} fusion_epochs={fusion_epochs}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
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]
|
||||
|
||||
# Holdout eval (if available)
|
||||
fu_hld_auc = nan
|
||||
if holdout_loader is not None:
|
||||
y_fu_h, p_fu_h = collect_probs_fused(fused, holdout_loader, device)
|
||||
fu_hld_auc = _score_arrays(y_fu_h, p_fu_h, num_classes)[1]
|
||||
|
||||
is_best_fused = not np.isnan(fu_auc) and fu_auc > best_fused_auc
|
||||
if is_best_fused:
|
||||
best_fused_auc = fu_auc
|
||||
best_fused_state = copy.deepcopy(fused.state_dict())
|
||||
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)
|
||||
|
||||
is_best_hld_fused = not np.isnan(fu_hld_auc) and fu_hld_auc > best_holdout_fused_auc
|
||||
if is_best_hld_fused:
|
||||
best_holdout_fused_auc = fu_hld_auc
|
||||
snap_holdout_fused = {"auc": fu_hld_auc, "acc": _score_arrays(y_fu_h, p_fu_h, num_classes)[0]}
|
||||
|
||||
if (fep + 1) % max(1, getattr(args, "log_every", 1)) == 0:
|
||||
print(
|
||||
f" [fold {fold+1}] fusion ep{fep+1:>3} "
|
||||
f"loss={fu_loss:.4f} train_acc={fu_acc:.4f} "
|
||||
f"val_auc={fu_auc:.4f} hld_auc={fu_hld_auc:.4f}"
|
||||
f"{' *' if is_best_fused else ''}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
if best_fused_state is not None:
|
||||
if args.save_checkpoints:
|
||||
torch.save(best_fused_state, fold_dir / "best_fused.pt")
|
||||
print(
|
||||
f" [fold {fold+1}] BEST "
|
||||
f"fused_head(acc={snap_fused.get('acc', nan):.4f},"
|
||||
f"auc={snap_fused.get('auc', nan):.4f}) "
|
||||
f"kappa={snap_fused.get('kappa', nan):.4f} "
|
||||
f"F1={snap_fused.get('macro_f1', nan):.4f} "
|
||||
f"ECE={snap_fused.get('ece', nan):.4f} "
|
||||
f"holdout_auc={snap_holdout_fused.get('auc', nan):.4f}",
|
||||
flush=True,
|
||||
)
|
||||
|
||||
if run_single:
|
||||
if tower_mode == "single":
|
||||
print(
|
||||
@@ -911,6 +1013,11 @@ class V2HyperTower:
|
||||
else:
|
||||
y_bi_best = p_bi_best = None
|
||||
|
||||
y_fu_best = p_fu_best = None
|
||||
if run_fused and best_fused_state is not None:
|
||||
fused.load_state_dict(best_fused_state)
|
||||
y_fu_best, p_fu_best = collect_probs_fused(fused, val_loader, device)
|
||||
|
||||
return FoldResult(
|
||||
mode=mode, fold=fold,
|
||||
best_epoch_single=best_epoch_single, best_epoch_bilat=best_epoch_bilat,
|
||||
@@ -953,10 +1060,23 @@ class V2HyperTower:
|
||||
holdout_n=len(holdout_bilat),
|
||||
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_holdout_auc=snap_holdout_fused.get("auc", nan),
|
||||
fused_holdout_acc=snap_holdout_fused.get("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,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -976,6 +1096,7 @@ class V2HyperTower:
|
||||
("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"]:
|
||||
@@ -990,6 +1111,7 @@ class V2HyperTower:
|
||||
("classic_holdout", "classic_holdout"),
|
||||
("ensemble_holdout", "ensemble_holdout"),
|
||||
("bilat_holdout", "bilat_holdout"),
|
||||
("fused_holdout", "fused_holdout"),
|
||||
]:
|
||||
sub = {}
|
||||
for m in ["auc", "acc"]:
|
||||
@@ -1003,6 +1125,7 @@ class V2HyperTower:
|
||||
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_fused_vs_ensemble", "ensemble_val", "fused_val"),
|
||||
]:
|
||||
delta = {}
|
||||
for m in ["auc", "f1", "kappa"]:
|
||||
@@ -1022,7 +1145,8 @@ class V2HyperTower:
|
||||
out["eval_note"] = (
|
||||
"classic=eye-level SingleEyeHT; "
|
||||
"ensemble=patient-level SingleEyeHT (OD+OS averaged); "
|
||||
"bilateral=patient-level BilateralHT"
|
||||
"bilateral=patient-level BilateralHT; "
|
||||
"fused=ensemble base + learned logit-level fusion head"
|
||||
)
|
||||
return out
|
||||
|
||||
@@ -1034,8 +1158,11 @@ class V2HyperTower:
|
||||
cv = s["classic_best_val"]
|
||||
ev = s["ensemble_best_val"]
|
||||
bv = s["bilat_best_val"]
|
||||
fv = s["fused_best_val"]
|
||||
d1 = s["delta_ensemble_vs_classic"]
|
||||
d2 = s["delta_bilat_vs_ensemble"]
|
||||
d3 = s["delta_fused_vs_ensemble"]
|
||||
has_fused = fv["auc_mean"] is not None
|
||||
|
||||
print(f"\n=== Summary [{mode}] — best-epoch val ===")
|
||||
print(f" {'':26s} {'AUC':>8} {'ACC':>8} {'Kappa':>8} {'F1-mac':>8} {'ECE':>8}")
|
||||
@@ -1043,6 +1170,8 @@ class V2HyperTower:
|
||||
rows = [("single (eye-lvl eval)", cv)]
|
||||
elif tower_mode == "ensemble":
|
||||
rows = [("ensemble (pat-lvl eval)", ev)]
|
||||
if has_fused:
|
||||
rows.append(("fused_head(pat-lvl eval)", fv))
|
||||
elif tower_mode == "bilateral":
|
||||
rows = [("bilateral (bilat eval)", bv)]
|
||||
else:
|
||||
@@ -1051,6 +1180,8 @@ class V2HyperTower:
|
||||
("ensemble (pat-lvl eval)", ev),
|
||||
("bilateral (bilat eval)", bv),
|
||||
]
|
||||
if has_fused:
|
||||
rows.append(("fused_head(pat-lvl eval)", fv))
|
||||
for label, d in rows:
|
||||
print(
|
||||
f" {label:26s} "
|
||||
@@ -1068,6 +1199,12 @@ class V2HyperTower:
|
||||
f"{f(d2['auc_mean']):>8} {'':>8} "
|
||||
f"{f(d2['kappa_mean']):>8} {f(d2['f1_mean']):>8}"
|
||||
)
|
||||
if has_fused and tower_mode in ("ensemble", None):
|
||||
print(
|
||||
f" {'Δ fused−ensemble':26s} "
|
||||
f"{f(d3['auc_mean']):>8} {'':>8} "
|
||||
f"{f(d3['kappa_mean']):>8} {f(d3['f1_mean']):>8}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user