v4 update
This commit is contained in:
+326
-4
@@ -29,7 +29,6 @@ from v3.classes.croppers import (
|
||||
build_image_preprocessor_from_args,
|
||||
)
|
||||
from v3.classes.image_loader import CachedImageLoader
|
||||
from v3.classes.dataset import _ClinicalView # noqa: F401
|
||||
from v3.classes.loader_factory import (
|
||||
build_balanced_sampler,
|
||||
filter_bilateral_samples,
|
||||
@@ -37,7 +36,12 @@ from v3.classes.loader_factory import (
|
||||
make_loader,
|
||||
)
|
||||
from v3.classes.metrics import _score_arrays, _svf, _tune_and_snap
|
||||
from v3.classes.models import (
|
||||
from v3.classes.bridges import Bridge
|
||||
from v3.classes.towerbase import train_towers_epoch, collect_probs_towers
|
||||
from v3.classes.image_towers import ImageTower
|
||||
from v3.classes.clinical_towers import ClinicalDataTower
|
||||
from v3.classes.geometry_towers import GeometryTower
|
||||
from v3.classes.hypertower_models import (
|
||||
BilateralHT,
|
||||
EmbeddingMLPEnsembleHT,
|
||||
FusedEnsembleHT,
|
||||
@@ -154,7 +158,7 @@ class V3HyperTower:
|
||||
ap.add_argument("--exclude-cols", nargs="*", default=[])
|
||||
ap.add_argument("--eval-mode", choices=["binary", "multiclass"], default="binary")
|
||||
ap.add_argument(
|
||||
"--tower-mode", choices=["single", "ensemble", "bilateral", "siamese", "classic"],
|
||||
"--hypertower-mode", choices=["single", "ensemble", "bilateral", "siamese", "classic"],
|
||||
default="ensemble",
|
||||
)
|
||||
ap.add_argument("--n-splits", type=int, default=5)
|
||||
@@ -252,6 +256,20 @@ class V3HyperTower:
|
||||
ap.add_argument("--geometry-source", default="gt", choices=["gt", "unet"],
|
||||
help="Source for geometry features: gt (GT contour annotations) or "
|
||||
"unet (U-Net segmentation). unet also requires --img-crop-weights.")
|
||||
ap.add_argument("--geometry-tower", action="store_true",
|
||||
help="Add a dedicated GeometryTower (disc/cup seg-map CNN) fused via the "
|
||||
"bridge alongside ImageTower and ClinicalDataTower. Requires "
|
||||
"--img-crop-manifest.")
|
||||
ap.add_argument("--geometry-tower-backbone", default="resnet18",
|
||||
choices=["resnet18", "resnet50", "efficientnet_b0"],
|
||||
help="SegCNN backbone for GeometryTower (default: resnet18).")
|
||||
ap.add_argument("--geometry-tower-in-channels", type=int, default=3, choices=[1, 3],
|
||||
help="1 = single label map; 3 = one-hot disc/rim/cup (default: 3).")
|
||||
ap.add_argument("--geometry-tower-frozen", action="store_true",
|
||||
help="Freeze GeometryTower backbone throughout training.")
|
||||
ap.add_argument("--geometry-tower-finetune-unet-epochs", type=int, default=0,
|
||||
help="Epochs to fine-tune the U-Net per fold before seg-map extraction "
|
||||
"(0 = disabled; only applies when --geometry-source unet).")
|
||||
return ap
|
||||
|
||||
def __init__(self, args) -> None:
|
||||
@@ -333,7 +351,7 @@ class V3HyperTower:
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
mode = args.eval_mode
|
||||
tower_mode = "single" if args.tower_mode == "classic" else args.tower_mode
|
||||
tower_mode = "single" if args.hypertower_mode == "classic" else args.hypertower_mode
|
||||
df_mode = self.data.df.copy()
|
||||
|
||||
if args.exclude_mixed_patients:
|
||||
@@ -482,6 +500,300 @@ class V3HyperTower:
|
||||
|
||||
return out_dir
|
||||
|
||||
def _run_fold_towers(
|
||||
self,
|
||||
*,
|
||||
fold: int,
|
||||
split,
|
||||
mode: str,
|
||||
data,
|
||||
num_classes: int,
|
||||
profile_eye,
|
||||
profile_patient,
|
||||
fold_dir: Path,
|
||||
pred_store,
|
||||
image_cache,
|
||||
):
|
||||
"""Modular TowerBase training path (used when --geometry-tower is set).
|
||||
|
||||
Builds [ImageTower, ClinicalDataTower, GeometryTower], runs the full
|
||||
fold lifecycle (prepare_fold → augment_samples → loader build → epoch
|
||||
loop → eval), and returns (FoldResult, FoldArtifacts) with metrics in
|
||||
the ensemble_val_* slots.
|
||||
"""
|
||||
args = self.args
|
||||
device = self.device
|
||||
nan = float("nan")
|
||||
|
||||
# ------------------------------------------------------------------ samples
|
||||
eye_train = filter_eye_samples(profile_eye.build_samples(df=split.train, clinical=data))
|
||||
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))
|
||||
bilat_test = filter_bilateral_samples(profile_patient.build_samples(
|
||||
df=split.test, clinical=data)) if split.test is not None else []
|
||||
|
||||
# Old --geometry-dim path still applies (injects geometry into clinical stream)
|
||||
if self.geometry_provider is not None:
|
||||
eye_train = self._augment_geometry(eye_train)
|
||||
bilat_train = self._augment_geometry(bilat_train)
|
||||
bilat_val = self._augment_geometry(bilat_val)
|
||||
bilat_test = self._augment_geometry(bilat_test)
|
||||
|
||||
if len(bilat_val) == 0:
|
||||
empty = FoldResult(
|
||||
mode=mode, fold=fold,
|
||||
best_epoch_single=0, best_epoch_bilat=0,
|
||||
classic_val_auc=nan, classic_val_acc=nan, classic_val_kappa=nan,
|
||||
classic_val_mcc=nan, classic_val_f1=nan, classic_val_recall=None,
|
||||
classic_val_ece=nan, classic_val_threshold=nan, classic_val_bias=None,
|
||||
classic_val_n=0,
|
||||
ensemble_val_auc=nan, ensemble_val_acc=nan, ensemble_val_kappa=nan,
|
||||
ensemble_val_mcc=nan, ensemble_val_f1=nan, ensemble_val_recall=None,
|
||||
ensemble_val_ece=nan, ensemble_val_threshold=nan, ensemble_val_bias=None,
|
||||
ensemble_val_n=0,
|
||||
bilat_val_auc=nan, bilat_val_acc=nan, bilat_val_kappa=nan,
|
||||
bilat_val_mcc=nan, bilat_val_f1=nan, bilat_val_recall=None,
|
||||
bilat_val_ece=nan, bilat_val_threshold=nan, bilat_val_bias=None,
|
||||
bilat_val_n=0,
|
||||
single_train_n=len(eye_train), bilat_train_n=len(bilat_train),
|
||||
)
|
||||
return empty, FoldArtifacts(
|
||||
y_true_classic=None, probs_classic=None,
|
||||
y_true_ensemble=None, probs_ensemble=None,
|
||||
y_true_bilat=None, probs_bilat=None,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ towers
|
||||
img_tower = ImageTower(
|
||||
backbone=args.backbone,
|
||||
freeze_ratio=args.freeze_ratio,
|
||||
augment=args.augment,
|
||||
use_se=getattr(args, "se_img_tower", False),
|
||||
)
|
||||
cd_tower = ClinicalDataTower(
|
||||
clinical_data=data,
|
||||
cd_hidden_dim=args.cd_hidden_dim,
|
||||
cd_dropout=getattr(args, "cd_dropout", 0.1),
|
||||
use_se=getattr(args, "se_cd_tower", False),
|
||||
)
|
||||
|
||||
geom_tower = GeometryTower(
|
||||
backbone=getattr(args, "geometry_tower_backbone", "resnet18"),
|
||||
in_channels=getattr(args, "geometry_tower_in_channels", 3),
|
||||
pretrained=not getattr(args, "no_pretrained", False),
|
||||
frozen=getattr(args, "geometry_tower_frozen", False),
|
||||
geometry_source=getattr(args, "geometry_source", "gt"),
|
||||
manifest_path=getattr(args, "img_crop_manifest", None),
|
||||
weights_path=getattr(args, "img_crop_weights", None),
|
||||
unet_normalize=getattr(args, "img_crop_normalize", "per_image"),
|
||||
unet_threshold=getattr(args, "img_crop_threshold", 0.5),
|
||||
finetune_unet_epochs=getattr(args, "geometry_tower_finetune_unet_epochs", 0),
|
||||
)
|
||||
|
||||
# GeometryTower.prepare_fold must run before augment_samples (precomputes seg maps)
|
||||
geom_tower.prepare_fold(
|
||||
eye_train=eye_train, bilat_train=bilat_train,
|
||||
bilat_val=bilat_val, bilat_test=bilat_test,
|
||||
image_preprocessor=self.image_preprocessor,
|
||||
image_cache=image_cache, device=device, args=args,
|
||||
)
|
||||
# Inject seg_map_1/seg_map_2 into all sample lists before loaders are built
|
||||
for sample_list in (eye_train, bilat_train, bilat_val, bilat_test):
|
||||
geom_tower.augment_samples(sample_list)
|
||||
|
||||
# Now ImageTower.prepare_fold sees augmented samples → loader includes seg maps
|
||||
img_tower.prepare_fold(
|
||||
eye_train=eye_train, bilat_train=bilat_train,
|
||||
bilat_val=bilat_val, bilat_test=bilat_test,
|
||||
image_preprocessor=self.image_preprocessor,
|
||||
image_cache=image_cache, device=device, args=args,
|
||||
)
|
||||
cd_tower.prepare_fold(
|
||||
eye_train=eye_train, bilat_train=bilat_train,
|
||||
bilat_val=bilat_val, bilat_test=bilat_test,
|
||||
image_preprocessor=self.image_preprocessor,
|
||||
image_cache=image_cache, device=device, args=args,
|
||||
)
|
||||
|
||||
towers = [img_tower, cd_tower, geom_tower]
|
||||
|
||||
# ------------------------------------------------------------------ bridge
|
||||
tower_dims = []
|
||||
for t in towers:
|
||||
tower_dims.extend(t.embed_dims)
|
||||
bridge = Bridge(
|
||||
tower_dims=tower_dims,
|
||||
num_classes=num_classes,
|
||||
fusion_dim=args.fusion_dim,
|
||||
mode=getattr(args, "bridge_mode", "fused"),
|
||||
dropout=getattr(args, "bridge_dropout", 0.5),
|
||||
)
|
||||
|
||||
# Move all nn.Modules to device
|
||||
for t in towers:
|
||||
if isinstance(t, torch.nn.Module):
|
||||
t.to(device)
|
||||
bridge.to(device)
|
||||
|
||||
# ------------------------------------------------------------------ loaders
|
||||
slots_patient = profile_patient.slot_descriptors()
|
||||
_persistent = args.num_workers > 0
|
||||
loader_kw = dict(
|
||||
batch_size=args.batch_size, num_workers=args.num_workers,
|
||||
image_cache=image_cache, persistent_workers=_persistent,
|
||||
)
|
||||
eval_transform = build_eval_transform(args.backbone)
|
||||
val_loader = make_loader(
|
||||
bilat_val, slots_patient,
|
||||
image_transform=eval_transform,
|
||||
image_preprocessor=self.image_preprocessor,
|
||||
shuffle=False, **loader_kw,
|
||||
)
|
||||
test_loader = None
|
||||
if bilat_test:
|
||||
test_loader = make_loader(
|
||||
bilat_test, slots_patient,
|
||||
image_transform=eval_transform,
|
||||
image_preprocessor=self.image_preprocessor,
|
||||
shuffle=False, **loader_kw,
|
||||
)
|
||||
|
||||
train_loader = img_tower.train_loader
|
||||
val_loader.dataset.prebuild_image_cache()
|
||||
if test_loader is not None:
|
||||
test_loader.dataset.prebuild_image_cache()
|
||||
|
||||
# ------------------------------------------------------------------ optimizer
|
||||
all_params = list(bridge.parameters())
|
||||
for t in towers:
|
||||
if isinstance(t, torch.nn.Module):
|
||||
all_params.extend(t.parameters())
|
||||
optimizer = torch.optim.AdamW(
|
||||
[p for p in all_params if p.requires_grad],
|
||||
lr=args.lr,
|
||||
weight_decay=getattr(args, "weight_decay", 1e-4),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ epoch loop
|
||||
global_warmup_tower = getattr(args, "warmup_tower_epochs", None)
|
||||
global_warmup_fused = getattr(args, "warmup_fused_epochs", None)
|
||||
warmup_cd = int(getattr(args, "warmup_cd_epochs", 0))
|
||||
warmup_tower = int(getattr(args, "single_warmup_tower_epochs", None) or global_warmup_tower or 2)
|
||||
warmup_fused = int(getattr(args, "single_warmup_fused_epochs", None) or global_warmup_fused or 2)
|
||||
main_epochs = int(args.epochs)
|
||||
|
||||
schedule = []
|
||||
if warmup_cd > 0: schedule.append(("cd_warmup", warmup_cd))
|
||||
if warmup_tower > 0: schedule.append(("tower_warmup", warmup_tower))
|
||||
if warmup_fused > 0: schedule.append(("fused_warmup", warmup_fused))
|
||||
schedule.append(("main", main_epochs))
|
||||
|
||||
best_val_auc = float("-inf")
|
||||
best_epoch = 0
|
||||
best_tower_states = None
|
||||
best_bridge_state = None
|
||||
epoch_idx = 0
|
||||
|
||||
for phase, n_epochs in schedule:
|
||||
for _ in range(n_epochs):
|
||||
for t in towers:
|
||||
if isinstance(t, torch.nn.Module):
|
||||
t.train()
|
||||
train_towers_epoch(
|
||||
towers, bridge, train_loader, optimizer, device,
|
||||
phase=phase,
|
||||
bcd_prob=getattr(args, "bcd_prob", 0.5),
|
||||
tower_loss_mode=getattr(args, "tower_loss_mode", "bcd"),
|
||||
)
|
||||
y_v, p_v = collect_probs_towers(towers, bridge, val_loader, device,
|
||||
tower_mode="ensemble")
|
||||
_, val_auc, _ = _score_arrays(y_v, p_v, num_classes)
|
||||
if not np.isnan(val_auc) and val_auc > best_val_auc:
|
||||
best_val_auc = val_auc
|
||||
best_epoch = epoch_idx
|
||||
best_tower_states = [
|
||||
t.state_dict() if isinstance(t, torch.nn.Module) else None
|
||||
for t in towers
|
||||
]
|
||||
best_bridge_state = bridge.state_dict()
|
||||
epoch_idx += 1
|
||||
|
||||
# Restore best
|
||||
if best_bridge_state is not None:
|
||||
bridge.load_state_dict(best_bridge_state)
|
||||
if best_tower_states is not None:
|
||||
for t, st in zip(towers, best_tower_states):
|
||||
if isinstance(t, torch.nn.Module) and st is not None:
|
||||
t.load_state_dict(st)
|
||||
|
||||
# ------------------------------------------------------------------ eval
|
||||
y_val, p_val = collect_probs_towers(towers, bridge, val_loader, device, tower_mode="ensemble")
|
||||
acc_val, auc_val, n_val = _score_arrays(y_val, p_val, num_classes)
|
||||
snap_val, _, thr_val, bias_val = _tune_and_snap(
|
||||
y_val, p_val, acc_val, num_classes, args, n_bins=10
|
||||
)
|
||||
|
||||
y_test = p_test = None
|
||||
test_auc = test_acc = nan
|
||||
test_n = 0
|
||||
if test_loader is not None:
|
||||
y_test, p_test = collect_probs_towers(towers, bridge, test_loader, device,
|
||||
tower_mode="ensemble")
|
||||
test_acc, test_auc, test_n = _score_arrays(y_test, p_test, num_classes)
|
||||
|
||||
result = FoldResult(
|
||||
mode=mode, fold=fold,
|
||||
best_epoch_single=best_epoch, best_epoch_bilat=0,
|
||||
classic_val_auc=nan, classic_val_acc=nan, classic_val_kappa=nan,
|
||||
classic_val_mcc=nan, classic_val_f1=nan, classic_val_recall=None,
|
||||
classic_val_ece=nan, classic_val_threshold=nan, classic_val_bias=None,
|
||||
classic_val_n=0,
|
||||
ensemble_val_auc=snap_val["auc"], ensemble_val_acc=snap_val["acc"],
|
||||
ensemble_val_kappa=snap_val["kappa"], ensemble_val_mcc=snap_val["mcc"],
|
||||
ensemble_val_f1=snap_val["macro_f1"],
|
||||
ensemble_val_recall=_sv(snap_val["per_class_recall"]),
|
||||
ensemble_val_ece=snap_val["ece"],
|
||||
ensemble_val_threshold=snap_val["threshold"],
|
||||
ensemble_val_bias=_svf(bias_val),
|
||||
ensemble_val_n=snap_val["n"],
|
||||
bilat_val_auc=nan, bilat_val_acc=nan, bilat_val_kappa=nan,
|
||||
bilat_val_mcc=nan, bilat_val_f1=nan, bilat_val_recall=None,
|
||||
bilat_val_ece=nan, bilat_val_threshold=nan, bilat_val_bias=None,
|
||||
bilat_val_n=0,
|
||||
ensemble_test_auc=test_auc, ensemble_test_acc=test_acc,
|
||||
test_n=test_n,
|
||||
single_train_n=len(eye_train), bilat_train_n=len(bilat_train),
|
||||
)
|
||||
artifacts = FoldArtifacts(
|
||||
y_true_classic=None, probs_classic=None,
|
||||
y_true_ensemble=y_val, probs_ensemble=p_val,
|
||||
y_true_bilat=None, probs_bilat=None,
|
||||
y_true_test=y_test, probs_test=p_test,
|
||||
)
|
||||
return result, artifacts
|
||||
|
||||
def _augment_geometry_slot(self, samples: list) -> list:
|
||||
"""Add geom_1/geom_2 keys to each sample dict (geometry tower mode).
|
||||
|
||||
Unlike _augment_geometry, this does NOT touch matrix_1/matrix_2 — the
|
||||
geometry vector lives in its own slot so ImageTower and ClinicalDataTower
|
||||
each receive only their own modality.
|
||||
"""
|
||||
if self.geometry_provider is None:
|
||||
return samples
|
||||
geom_dim = int(getattr(self.args, "geometry_dim", 0)) or 5
|
||||
for s in samples:
|
||||
for img_slot, geom_slot in (("image_1", "geom_1"), ("image_2", "geom_2")):
|
||||
img_path = s.get(img_slot)
|
||||
if img_path is None:
|
||||
continue
|
||||
vec = self.geometry_provider.geometry_for_image(img_path)
|
||||
if vec is not None and len(vec) >= geom_dim:
|
||||
s[geom_slot] = vec[:geom_dim].astype(np.float32)
|
||||
else:
|
||||
s[geom_slot] = np.zeros(geom_dim, dtype=np.float32)
|
||||
return samples
|
||||
|
||||
def _augment_geometry(self, samples: list) -> list:
|
||||
"""Append geometry features to matrix_1/matrix_2 in each sample dict."""
|
||||
if self.geometry_provider is None:
|
||||
@@ -517,6 +829,16 @@ class V3HyperTower:
|
||||
image_cache,
|
||||
):
|
||||
args = self.args
|
||||
|
||||
# Modular tower path — bypasses the legacy single/bilat/siamese code entirely
|
||||
if getattr(args, "geometry_tower", False):
|
||||
return self._run_fold_towers(
|
||||
fold=fold, split=split, mode=mode, data=data,
|
||||
num_classes=num_classes, profile_eye=profile_eye,
|
||||
profile_patient=profile_patient, fold_dir=fold_dir,
|
||||
pred_store=pred_store, image_cache=image_cache,
|
||||
)
|
||||
|
||||
device = self.device
|
||||
image_preprocessor = self.image_preprocessor
|
||||
nan = float("nan")
|
||||
|
||||
Reference in New Issue
Block a user