v4 update

This commit is contained in:
rpotter6298
2026-04-20 18:01:31 +02:00
parent 13290575d5
commit 4dea45df78
71 changed files with 8316 additions and 4112 deletions
+326 -4
View File
@@ -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")