From 13ad32683ffaec455bd8ec911086d931e19ac4f5 Mon Sep 17 00:00:00 2001 From: rpotter6298 Date: Wed, 4 Mar 2026 11:17:10 +0100 Subject: [PATCH] predictions class added --- classes/v2/predictions.py | 194 +++++++++++++++++++++++++ classes/v2/v2_hypertower.py | 128 +++++++++++++++- scripts/main/v2/run_binary_v2.2.sh | 33 ++--- scripts/main/v2/run_multiclass_v2.2.sh | 32 ++-- 4 files changed, 338 insertions(+), 49 deletions(-) create mode 100644 classes/v2/predictions.py diff --git a/classes/v2/predictions.py b/classes/v2/predictions.py new file mode 100644 index 0000000..009ec39 --- /dev/null +++ b/classes/v2/predictions.py @@ -0,0 +1,194 @@ +"""PredictionStore — unified per-epoch prediction tensor across all folds. + +Tensor shape: (n_folds, n_epochs, n_samples, n_heads, n_classes) + +The meaning of "sample" depends on tower_mode: + single — each eye is a sample; sample_ids like "5OD", "14OS" + ensemble — each patient is a sample; sample_ids like "5", "14" + fused — same as ensemble + bilateral— same as ensemble + +Head names by mode: + single : ["fused", "img", "md"] + ensemble : ["od_fused", "od_img", "od_md", "os_fused", "os_img", "os_md"] + fused : ["od_fused", "od_img", "od_md", "os_fused", "os_img", "os_md", "bilat_fused"] + bilateral : ["fused", "img_joint", "md_joint"] +""" +from __future__ import annotations + +from pathlib import Path +from typing import Sequence + +import numpy as np + + +def head_names_for_mode(tower_mode: str, *, fused_head: bool = False) -> list[str]: + """Return canonical head name list for a given tower_mode.""" + if tower_mode in ("single", "classic"): + return ["fused", "img", "md"] + if tower_mode == "ensemble": + names = ["od_fused", "od_img", "od_md", "os_fused", "os_img", "os_md"] + return names + ["bilat_fused"] if fused_head else names + if tower_mode == "bilateral": + return ["fused", "img_joint", "md_joint"] + raise ValueError(f"Unknown tower_mode: {tower_mode!r}") + + +class PredictionStore: + """ + Stores per-epoch predictions for every sample, head, and fold in one tensor. + + Usage + ----- + # Build once before the fold loop: + store = PredictionStore( + sample_ids=all_eye_or_patient_ids, + y_true=all_labels, + head_names=head_names_for_mode(tower_mode, fused_head=args.fused_head), + n_folds=n_folds, + n_epochs=total_epochs, + n_classes=num_classes, + ) + + # Inside each epoch, after collecting probs: + store.record(fold, epoch, patient_ids_batch, "od_fused", probs_od) + store.set_split(fold, train_ids, "train") + store.set_split(fold, val_ids, "val") + + # After all folds: + store.save(run_dir / "predictions.npz") + + # Load and query: + store = PredictionStore.load("predictions.npz") + store.query("5", "od_fused", fold=0) # → (n_epochs, n_classes) + store.query("5", "od_fused") # → (n_folds, n_epochs, n_classes) + store.get_split("5", fold=0) # → "train" + """ + + def __init__( + self, + sample_ids: Sequence[str], + y_true: Sequence[int], + head_names: Sequence[str], + n_folds: int, + n_epochs: int, + n_classes: int, + ): + self.sample_ids = np.array(sample_ids, dtype=object) + self.y_true = np.array(y_true, dtype=np.int64) + self.head_names = np.array(head_names, dtype=object) + self.n_folds = n_folds + self.n_epochs = n_epochs + self.n_classes = n_classes + + n_samples = len(self.sample_ids) + n_heads = len(self.head_names) + + self.probs = np.full( + (n_folds, n_epochs, n_samples, n_heads, n_classes), + fill_value=np.nan, + dtype=np.float32, + ) + self.split = np.full((n_folds, n_samples), fill_value="", dtype=object) + + self._sid_index: dict[str, int] = {str(s): i for i, s in enumerate(self.sample_ids)} + self._head_index: dict[str, int] = {str(h): i for i, h in enumerate(self.head_names)} + + # ------------------------------------------------------------------ + # Writing + # ------------------------------------------------------------------ + + def record( + self, + fold: int, + epoch: int, + sample_ids: Sequence[str], + head_name: str, + probs: np.ndarray, + ) -> None: + """Record a batch of predictions for one head. + + Args: + fold: 0-indexed fold number + epoch: 0-indexed epoch number + sample_ids: sequence of sample ID strings (length B) + head_name: which head — must be in self.head_names + probs: (B, n_classes) probability array + """ + head_idx = self._head_index.get(head_name) + if head_idx is None: + return # head not active in this mode — skip silently + for i, sid in enumerate(sample_ids): + s_idx = self._sid_index.get(str(sid)) + if s_idx is not None: + self.probs[fold, epoch, s_idx, head_idx, :] = probs[i] + + def set_split( + self, + fold: int, + sample_ids: Sequence[str], + label: str, + ) -> None: + """Label a group of samples as 'train', 'val', or 'holdout' for a fold.""" + for sid in sample_ids: + s_idx = self._sid_index.get(str(sid)) + if s_idx is not None: + self.split[fold, s_idx] = label + + # ------------------------------------------------------------------ + # Querying + # ------------------------------------------------------------------ + + def query( + self, + sample_id: str, + head_name: str, + fold: int | None = None, + ) -> np.ndarray: + """Return epoch-level predictions for one sample + head. + + Returns: + fold=None → (n_folds, n_epochs, n_classes) + fold=int → (n_epochs, n_classes) + """ + s_idx = self._sid_index[str(sample_id)] + head_idx = self._head_index[str(head_name)] + if fold is None: + return self.probs[:, :, s_idx, head_idx, :] + return self.probs[fold, :, s_idx, head_idx, :] + + def get_split(self, sample_id: str, fold: int) -> str: + """Return the split label ('train'/'val'/'holdout') for a sample in a fold.""" + s_idx = self._sid_index[str(sample_id)] + return str(self.split[fold, s_idx]) + + # ------------------------------------------------------------------ + # Persistence + # ------------------------------------------------------------------ + + def save(self, path: str | Path) -> None: + np.savez_compressed( + path, + probs=self.probs, + split=self.split, + sample_ids=self.sample_ids, + y_true=self.y_true, + head_names=self.head_names, + ) + + @classmethod + def load(cls, path: str | Path) -> "PredictionStore": + data = np.load(path, allow_pickle=True) + probs = data["probs"] + n_folds, n_epochs, _, _, n_classes = probs.shape + store = cls( + sample_ids=data["sample_ids"].tolist(), + y_true=data["y_true"], + head_names=data["head_names"].tolist(), + n_folds=n_folds, + n_epochs=n_epochs, + n_classes=n_classes, + ) + store.probs = probs + store.split = data["split"] + return store diff --git a/classes/v2/v2_hypertower.py b/classes/v2/v2_hypertower.py index 81baf80..ba54d0c 100644 --- a/classes/v2/v2_hypertower.py +++ b/classes/v2/v2_hypertower.py @@ -48,6 +48,7 @@ from classes.v2.models import ( train_single_epoch, ) from classes.v2.papila_builders import build_papila_data +from classes.v2.predictions import PredictionStore, head_names_for_mode from classes.v2.profiles import build_papila_profile from classes.v2.results import FoldArtifacts, FoldResult, _f, _nan, _sv from classes.v2.split_manager import PatientFirstSplitManager @@ -407,6 +408,46 @@ class V2HyperTower: patient_col="Patient ID", label_col=args.label_col, sample_mode="patient" ) + # ---- PredictionStore — build once before fold loop --------------- + fused_head = getattr(args, "fused_head", False) + _head_names = head_names_for_mode(tower_mode, fused_head=fused_head) + fusion_epochs = int(getattr(args, "fusion_epochs", 10)) if fused_head else 0 + _global_warmup_tower = getattr(args, "warmup_tower_epochs", None) + _global_warmup_fused = getattr(args, "warmup_fused_epochs", None) + _warmup_tower = ( + int(args.single_warmup_tower_epochs) + if getattr(args, "single_warmup_tower_epochs", None) is not None + else int(_global_warmup_tower) if _global_warmup_tower is not None else 2 + ) + _warmup_fused = ( + int(args.single_warmup_fused_epochs) + 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 + + # sample IDs depend on mode: single uses eye IDs, others use patient IDs + if tower_mode in ("single", "classic"): + _sample_ids = [ + f"{row['Patient ID']}{row['eyeID']}" + for _, row in df_mode.iterrows() + ] + _y_true = df_mode[args.label_col].tolist() + else: + # one row per patient (deduplicate — take first occurrence per patient) + _pat_df = df_mode.drop_duplicates(subset="Patient ID") + _sample_ids = _pat_df["Patient ID"].astype(str).tolist() + _y_true = _pat_df[args.label_col].tolist() + + pred_store = PredictionStore( + sample_ids=_sample_ids, + y_true=_y_true, + head_names=_head_names, + n_folds=n_folds, + n_epochs=_total_epochs, + n_classes=num_classes, + ) + for fold in range(n_folds): seed_everything(args.seed + fold * 100) fold_dir = tm_dir / f"fold{fold}" @@ -423,6 +464,7 @@ class V2HyperTower: profile_patient=profile_patient, fold_dir=fold_dir, tower_mode=tower_mode, + pred_store=pred_store, ) fold_results.append(result) if artifacts.y_true_ensemble is not None: @@ -549,6 +591,7 @@ class V2HyperTower: "mode_summary": summary, } (tm_dir / "summary.json").write_text(json.dumps(mode_summary, indent=2), encoding="utf-8") + pred_store.save(tm_dir / "predictions.npz") root_summary_path = out_dir / "summary.json" if root_summary_path.exists(): @@ -583,6 +626,7 @@ class V2HyperTower: profile_patient, fold_dir: Path, tower_mode: str, + pred_store: "PredictionStore | None" = None, ) -> tuple[FoldResult, FoldArtifacts]: args = self.args device = self.device @@ -631,6 +675,18 @@ class V2HyperTower: bilat_train = filter_bilateral_samples(profile_patient.build_samples(df=split.train, clinical=data)) bilat_val = filter_bilateral_samples(profile_patient.build_samples(df=split.val, clinical=data)) + # Register split labels in the prediction store + if pred_store is not None: + if tower_mode in ("single", "classic"): + # eye-level IDs: "{patient_id}{eyeID}" + train_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in eye_train] + val_sids = [f"{s['id_1']}{s.get('eye_id_1','')}" for s in bilat_val] + else: + train_sids = [str(s["id_1"]) for s in bilat_train] + val_sids = [str(s["id_1"]) for s in bilat_val] + pred_store.set_split(fold, train_sids, "train") + pred_store.set_split(fold, val_sids, "val") + if len(bilat_val) == 0: print(f" [fold {fold+1}] WARNING: no bilateral val samples; skipping fold.", flush=True) empty = FoldResult( @@ -751,6 +807,8 @@ class V2HyperTower: **loader_kw, ) print(f" [fold {fold+1}] holdout_n={len(holdout_bilat)} (bilateral patients)", flush=True) + if pred_store is not None: + pred_store.set_split(fold, [str(s["id_1"]) for s in holdout_bilat], "holdout") 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 @@ -918,7 +976,7 @@ class V2HyperTower: _, cl_auc_img, _ = _score_arrays(y_cl, p_cl_img, num_classes) _, cl_auc_md, _ = _score_arrays(y_cl, p_cl_md, num_classes) y_en = np.array([], dtype=np.int64) - p_en = np.zeros((0, 0), dtype=np.float32) + p_en = p_en_img = p_en_md = np.zeros((0, num_classes), dtype=np.float32) en_acc = en_auc = nan en_n = 0 en_acc_img = en_acc_md = en_auc_img = en_auc_md = nan @@ -939,13 +997,14 @@ class V2HyperTower: _, en_auc_img, _ = _score_arrays(y_en, p_en_img, num_classes) _, en_auc_md, _ = _score_arrays(y_en, p_en_md, num_classes) y_cl = np.array([], dtype=np.int64) - p_cl = np.zeros((0, 0), dtype=np.float32) + p_cl = p_cl_img = p_cl_md = np.zeros((0, num_classes), dtype=np.float32) cl_acc = cl_auc = nan cl_n = 0 cl_acc_img = cl_acc_md = cl_auc_img = cl_auc_md = nan else: y_cl = y_en = np.array([], dtype=np.int64) - p_cl = p_en = np.zeros((0, 0), dtype=np.float32) + p_cl = p_cl_img = p_cl_md = np.zeros((0, num_classes), dtype=np.float32) + p_en = p_en_img = p_en_md = np.zeros((0, num_classes), dtype=np.float32) cl_acc = cl_auc = en_acc = en_auc = nan cl_n = en_n = 0 cl_acc_img = cl_acc_md = en_acc_img = en_acc_md = nan @@ -1052,6 +1111,23 @@ class V2HyperTower: _epoch_train_pm.append(p_tr_m) _epoch_train_ids.append(tr_ids) _epoch_train_y.append(y_tr) + # record into PredictionStore + if pred_store is not None: + if tower_mode in ("single", "classic"): + pred_store.record(fold, epoch, tr_ids, "fused", p_tr_f) + pred_store.record(fold, epoch, tr_ids, "img", p_tr_i) + pred_store.record(fold, epoch, tr_ids, "md", p_tr_m) + else: # ensemble: separate OD and OS by eye suffix + od_mask = np.array([str(i).endswith("OD") for i in tr_ids]) + os_mask = ~od_mask + od_pids = [str(i)[:-2] for i in tr_ids[od_mask]] + os_pids = [str(i)[:-2] for i in tr_ids[os_mask]] + pred_store.record(fold, epoch, od_pids, "od_fused", p_tr_f[od_mask]) + pred_store.record(fold, epoch, od_pids, "od_img", p_tr_i[od_mask]) + pred_store.record(fold, epoch, od_pids, "od_md", p_tr_m[od_mask]) + pred_store.record(fold, epoch, os_pids, "os_fused", p_tr_f[os_mask]) + pred_store.record(fold, epoch, os_pids, "os_img", p_tr_i[os_mask]) + pred_store.record(fold, epoch, os_pids, "os_md", p_tr_m[os_mask]) # accumulate val for npy tensors if run_single and tower_mode == "ensemble" and y_en.size: @@ -1073,6 +1149,20 @@ class V2HyperTower: _epoch_val_pm_os.append(p_cl_md) _epoch_val_y.append(y_cl) + # record val into PredictionStore + if pred_store is not None: + if run_single and tower_mode == "ensemble" and y_en.size: + pred_store.record(fold, epoch, _en_pat_ids, "od_fused", _p_en_f_od) + pred_store.record(fold, epoch, _en_pat_ids, "od_img", _p_en_i_od) + pred_store.record(fold, epoch, _en_pat_ids, "od_md", _p_en_m_od) + pred_store.record(fold, epoch, _en_pat_ids, "os_fused", _p_en_f_os) + pred_store.record(fold, epoch, _en_pat_ids, "os_img", _p_en_i_os) + pred_store.record(fold, epoch, _en_pat_ids, "os_md", _p_en_m_os) + elif run_single and tower_mode == "single" and y_cl.size: + # val in single mode: collect_probs_single_components(aggregate_patient=False) + # returns interleaved [all_OD, all_OS] per batch — IDs not tracked here yet + pass # single-mode val IDs not currently available; train IDs are sufficient + # Best-epoch checks (restricted to main phase). target_single_auc = cl_auc if tower_mode == "single" else en_auc target_holdout_single_auc = cl_auc_h if tower_mode == "single" else en_auc_h @@ -1197,11 +1287,33 @@ class V2HyperTower: 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_total = single_warmup_tower + elif phase_single == "fused_warmup": + single_phase_epoch = epoch - single_warmup_tower + 1 + single_phase_total = single_warmup_fused + else: + single_phase_epoch = main_epoch_single + single_phase_total = main_epochs + + if phase_bilat == "tower_warmup": + bilat_phase_epoch = epoch + 1 + bilat_phase_total = bilat_warmup_tower + elif phase_bilat == "fused_warmup": + bilat_phase_epoch = epoch - bilat_warmup_tower + 1 + bilat_phase_total = bilat_warmup_fused + else: + bilat_phase_epoch = main_epoch_bilat + bilat_phase_total = main_epochs + if run_single: if tower_mode == "single": msg = ( f" ep {epoch+1:>3}/{total_epochs} " - f"[single:{phase_single} {main_epoch_single}/{main_epochs}] " + f"[single:{phase_single} {single_phase_epoch}/{single_phase_total}] " f"fused(acc={cl_acc:.4f},auc={cl_auc:.4f}) " f"img(acc={cl_acc_img:.4f},auc={cl_auc_img:.4f}) " f"md(acc={cl_acc_md:.4f},auc={cl_auc_md:.4f}) " @@ -1211,7 +1323,7 @@ class V2HyperTower: else: msg = ( f" ep {epoch+1:>3}/{total_epochs} " - f"[single:{phase_single} {main_epoch_single}/{main_epochs}] " + f"[single:{phase_single} {single_phase_epoch}/{single_phase_total}] " f"fused(acc={en_acc:.4f},auc={en_auc:.4f}) " f"img(acc={en_acc_img:.4f},auc={en_auc_img:.4f}) " f"md(acc={en_acc_md:.4f},auc={en_auc_md:.4f}) " @@ -1221,7 +1333,7 @@ class V2HyperTower: else: msg = ( f" ep {epoch+1:>3}/{total_epochs} " - f"[bilat:{phase_bilat} {main_epoch_bilat}/{main_epochs}] " + f"[bilat:{phase_bilat} {bilat_phase_epoch}/{bilat_phase_total}] " f"fused(acc={bi_acc:.4f},auc={bi_auc:.4f}) " f"img(acc={bi_acc_img:.4f},auc={bi_auc_img:.4f}) " f"md(acc={bi_acc_md:.4f},auc={bi_auc_md:.4f}) " @@ -1285,10 +1397,14 @@ class V2HyperTower: flush=True, ) + _val_pids_for_store = [str(s["id_1"]) for s in bilat_val] 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] + if pred_store is not None and y_fu.size: + _store_ep = total_single_epochs + fep + pred_store.record(fold, _store_ep, _val_pids_for_store, "bilat_fused", p_fu) # Holdout eval (if available) fu_hld_auc = nan diff --git a/scripts/main/v2/run_binary_v2.2.sh b/scripts/main/v2/run_binary_v2.2.sh index 2b18228..1a1cc2e 100755 --- a/scripts/main/v2/run_binary_v2.2.sh +++ b/scripts/main/v2/run_binary_v2.2.sh @@ -1,9 +1,9 @@ #!/usr/bin/env bash set -euo pipefail -# Binary runs v2.2 (6 total): -# UNet crop: single | ensemble | fused head -# GT crop: single | ensemble | fused head +# Binary runs v2.2 (4 total): +# UNet crop: single | fused head +# GT crop: single | fused head ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../.." && pwd)" cd "$ROOT_DIR" @@ -17,6 +17,8 @@ COMMON=( --batch-size 8 --backbone refugelike --eval-mode binary + --single-warmup-tower-epochs 4 + --single-warmup-fused-epochs 4 --img-crop-manifest "$MANIFEST" ) @@ -30,45 +32,32 @@ GT_CROP=( # ── UNet crop ──────────────────────────────────────────────────────────────── -echo "[1/6] UNet crop — binary, single..." +echo "[1/4] UNet crop — binary, single..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${UNET_CROP[@]}" \ --tower-mode single \ --run-name v2.2_single_binary_unet_40ep_5fold -echo "[2/6] UNet crop — binary, ensemble..." +echo "[2/4] UNet crop — binary, ensemble + fused head..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${UNET_CROP[@]}" \ --tower-mode ensemble \ - --run-name v2.2_ensemble_binary_unet_40ep_5fold - -echo "[3/6] UNet crop — binary, ensemble + fused head..." -python3 scripts/main/v2/run_multifold_v2.py \ - "${COMMON[@]}" "${UNET_CROP[@]}" \ - --tower-mode ensemble \ - --fused-head --fusion-epochs 10 \ + --fused-head --fusion-epochs 20 \ --run-name v2.2_fused_binary_unet_40ep_5fold # ── GT crop ────────────────────────────────────────────────────────────────── -echo "[4/6] GT crop — binary, single..." +echo "[3/4] GT crop — binary, single..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${GT_CROP[@]}" \ --tower-mode single \ --run-name v2.2_single_binary_gt_40ep_5fold -echo "[5/6] GT crop — binary, ensemble..." +echo "[4/4] GT crop — binary, ensemble + fused head..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${GT_CROP[@]}" \ --tower-mode ensemble \ - --run-name v2.2_ensemble_binary_gt_40ep_5fold - -echo "[6/6] GT crop — binary, ensemble + fused head..." -python3 scripts/main/v2/run_multifold_v2.py \ - "${COMMON[@]}" "${GT_CROP[@]}" \ - --tower-mode ensemble \ - --img-crop-gt \ - --fused-head --fusion-epochs 10 \ + --fused-head --fusion-epochs 20 \ --run-name v2.2_fused_binary_gt_40ep_5fold echo "Binary v2.2 runs complete." diff --git a/scripts/main/v2/run_multiclass_v2.2.sh b/scripts/main/v2/run_multiclass_v2.2.sh index 64a8be1..8e0302e 100755 --- a/scripts/main/v2/run_multiclass_v2.2.sh +++ b/scripts/main/v2/run_multiclass_v2.2.sh @@ -1,9 +1,9 @@ #!/usr/bin/env bash set -euo pipefail -# Multiclass runs v2.2 (6 total): -# UNet crop: single | ensemble | fused head -# GT crop: single | ensemble | fused head +# Multiclass runs v2.2 (4 total): +# UNet crop: single | fused head +# GT crop: single | fused head ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../.." && pwd)" cd "$ROOT_DIR" @@ -17,6 +17,8 @@ COMMON=( --batch-size 8 --backbone refugelike --eval-mode multiclass + --single-warmup-tower-epochs 4 + --single-warmup-fused-epochs 4 --img-crop-manifest "$MANIFEST" ) @@ -30,44 +32,32 @@ GT_CROP=( # ── UNet crop ──────────────────────────────────────────────────────────────── -echo "[1/6] UNet crop — multiclass, single..." +echo "[1/4] UNet crop — multiclass, single..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${UNET_CROP[@]}" \ --tower-mode single \ --run-name v2.2_single_multiclass_unet_40ep_5fold -echo "[2/6] UNet crop — multiclass, ensemble..." +echo "[2/4] UNet crop — multiclass, ensemble + fused head..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${UNET_CROP[@]}" \ --tower-mode ensemble \ - --run-name v2.2_ensemble_multiclass_unet_40ep_5fold - -echo "[3/6] UNet crop — multiclass, ensemble + fused head..." -python3 scripts/main/v2/run_multifold_v2.py \ - "${COMMON[@]}" "${UNET_CROP[@]}" \ - --tower-mode ensemble \ - --fused-head --fusion-epochs 10 \ + --fused-head --fusion-epochs 20 \ --run-name v2.2_fused_multiclass_unet_40ep_5fold # ── GT crop ────────────────────────────────────────────────────────────────── -echo "[4/6] GT crop — multiclass, single..." +echo "[3/4] GT crop — multiclass, single..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${GT_CROP[@]}" \ --tower-mode single \ --run-name v2.2_single_multiclass_gt_40ep_5fold -echo "[5/6] GT crop — multiclass, ensemble..." +echo "[4/4] GT crop — multiclass, ensemble + fused head..." python3 scripts/main/v2/run_multifold_v2.py \ "${COMMON[@]}" "${GT_CROP[@]}" \ --tower-mode ensemble \ - --run-name v2.2_ensemble_multiclass_gt_40ep_5fold - -echo "[6/6] GT crop — multiclass, ensemble + fused head..." -python3 scripts/main/v2/run_multifold_v2.py \ - "${COMMON[@]}" "${GT_CROP[@]}" \ - --tower-mode ensemble \ - --fused-head --fusion-epochs 10 \ + --fused-head --fusion-epochs 20 \ --run-name v2.2_fused_multiclass_gt_40ep_5fold echo "Multiclass v2.2 runs complete."