added fused classifier head

This commit is contained in:
rpotter6298
2026-03-02 06:44:22 +01:00
parent fd1452c187
commit 080b5999fa
11 changed files with 583 additions and 21 deletions
+142 -5
View File
@@ -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" {'Δ fusedensemble':26s} "
f"{f(d3['auc_mean']):>8} {'':>8} "
f"{f(d3['kappa_mean']):>8} {f(d3['f1_mean']):>8}"
)
# ---------------------------------------------------------------------------