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
+61 -31
View File
@@ -21,14 +21,6 @@ from .slot_dataset import SlotDataset, slot_collate
from .papila_data import PapilaData
from .papila_builders import build_papila_data
from .data_bundle import DataBundle
from .dataset import ClinicalDataset
from .config_builder import (
ConfigAssembly,
assemble_config,
load_config,
resolve_imports,
)
from .filters import RegexFilter, ColumnFilter, apply_regex_filters, apply_column_filters
from .transforms import (
ImageTransformConfig,
backbone_transform_config,
@@ -43,11 +35,35 @@ from .transforms import (
TRANSFORM_REGISTRY,
build_transform_chain,
)
from .model_builder import V2ModelBundle, build_model_bundle
from .towers import ImageTower, ClinicalTower, SiameseImageTower, build_backbone
from .bridges import Bridge, VoteBridge
from .models import SingleEyeHT, BilateralHT
from .v2_hypertower import V2HyperTower, V2ModeComparisonOps, V2ModeComparator
from .towerbase import TowerBase, build_backbone, train_towers_epoch, collect_probs_towers
from .image_towers import ImageEncoder, SiameseImageTower, ImageTower
from .clinical_towers import ClinicalEncoder, ClinicalDataTower
from .geometry_towers import GeometryTower
from .hypertower_models import (
SingleEyeHT,
BilateralHT,
SiameseHT,
FusedEnsembleHT,
LogitMLPEnsembleHT,
EmbeddingMLPEnsembleHT,
NTowerHT,
NLateralHT,
MonoTowerHT,
train_single_epoch,
train_bilateral_epoch,
train_siamese_epoch,
train_fusion_epoch,
train_ntower_epoch,
train_mono_epoch,
collect_probs_classic,
collect_probs_ensemble,
collect_probs_bilateral,
collect_probs_siamese,
collect_probs_ntower,
collect_probs_mono,
V2ModeComparisonOps,
)
from .bridges import Bridge, HTClassifier, HyperBridge, VoteBridge
from .hypertower_logger import HypertowerLogger
__all__ = [
@@ -66,18 +82,9 @@ __all__ = [
"PapilaData",
"build_papila_data",
"DataBundle",
"ClinicalDataset",
"SlotLoaderFactory",
"SlotDataset",
"slot_collate",
"ConfigAssembly",
"assemble_config",
"load_config",
"resolve_imports",
"RegexFilter",
"ColumnFilter",
"apply_regex_filters",
"apply_column_filters",
"ImageTransformConfig",
"backbone_transform_config",
"build_backbone_transform",
@@ -90,18 +97,41 @@ __all__ = [
"UnetMaskProvider",
"TRANSFORM_REGISTRY",
"build_transform_chain",
"V2ModelBundle",
"build_model_bundle",
"ImageTower",
"ClinicalTower",
"SiameseImageTower",
"TowerBase",
"build_backbone",
"Bridge",
"VoteBridge",
"train_towers_epoch",
"collect_probs_towers",
"ImageEncoder",
"SiameseImageTower",
"ImageTower",
"ClinicalEncoder",
"ClinicalDataTower",
"GeometryTower",
"SingleEyeHT",
"BilateralHT",
"V2HyperTower",
"SiameseHT",
"FusedEnsembleHT",
"LogitMLPEnsembleHT",
"EmbeddingMLPEnsembleHT",
"NTowerHT",
"NLateralHT",
"MonoTowerHT",
"train_single_epoch",
"train_bilateral_epoch",
"train_siamese_epoch",
"train_fusion_epoch",
"train_ntower_epoch",
"train_mono_epoch",
"collect_probs_classic",
"collect_probs_ensemble",
"collect_probs_bilateral",
"collect_probs_siamese",
"collect_probs_ntower",
"collect_probs_mono",
"Bridge",
"HTClassifier",
"HyperBridge",
"VoteBridge",
"V2ModeComparisonOps",
"V2ModeComparator",
"HypertowerLogger",
]
+254 -44
View File
@@ -1,19 +1,64 @@
from __future__ import annotations
import torch
import torch.nn as nn
from v3.classes.SE_attention import SEBlock, SEGateLogger
# ---------------------------------------------------------------------------
# HTClassifier — standalone classification head
# ---------------------------------------------------------------------------
class HTClassifier(nn.Module):
"""Minimal classification head: ReLU → Dropout → Linear(in_dim → num_classes).
Used as the output stage of Bridge, HyperBridge, and any vehicle that needs
a reusable, identifiable classifier type.
"""
def __init__(self, in_dim: int, num_classes: int, dropout: float = 0.5):
super().__init__()
self.head = nn.Sequential(
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(in_dim, num_classes),
)
def forward(self, z: torch.Tensor) -> torch.Tensor:
return self.head(z)
# ---------------------------------------------------------------------------
# Bridge — N-tower fusion
# ---------------------------------------------------------------------------
class Bridge(nn.Module):
"""
N-tower fusion bridge.
Takes a list of tower embeddings, projects each to a common ``fusion_dim``,
element-wise multiplies all projections, optionally applies an SE gate, then
classifies the fused representation via an HTClassifier.
Each tower also gets an auxiliary classification head (used for BCD training).
Construction
------------
``tower_dims`` is an ordered list of embedding dimensionalities — one entry per
embedding slot that will be passed to ``fuse()`` or ``forward()``.
Tower slots are accessed by index: ``W[i]``, ``ln[i]``, ``aux_heads[i]``.
The bridge has no knowledge of what modality each slot carries.
"""
def __init__(
self,
img_dim,
meta_dim,
num_classes,
fusion_dim=256,
mode="fused",
tower_dims: list[int],
num_classes: int,
fusion_dim: int = 256,
mode: str = "fused",
dropout: float = 0.5,
use_se: bool = True,
se_reduction: int = 16,
@@ -22,29 +67,30 @@ class Bridge(nn.Module):
super().__init__()
self.mode = mode
self.use_se = use_se
self.tower_dims = list(tower_dims)
# project towers to equal width
self.W_img = nn.Linear(img_dim, fusion_dim)
self.W_md = nn.Linear(meta_dim, fusion_dim)
# Per-tower projection heads: each projects dim_i → fusion_dim
self.W = nn.ModuleList([nn.Linear(d, fusion_dim) for d in tower_dims])
self.ln = nn.ModuleList(
[nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity()
for _ in tower_dims]
)
# optional: layernorm before SE
self.ln_img = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity()
self.ln_md = nn.LayerNorm(fusion_dim) if se_pre_norm else nn.Identity()
# Per-tower auxiliary classifiers (for BCD training)
self.aux_heads = nn.ModuleList([nn.Linear(d, num_classes) for d in tower_dims])
# SE gate on the fused vector
self.se = SEBlock(fusion_dim, reduction=se_reduction, residual=True) if use_se else None
self.se_log = SEGateLogger(enabled=use_se, track_channels=False, dim=fusion_dim)
# heads
self.classifier_fused = nn.Sequential(
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(fusion_dim, num_classes),
)
self.classifier_img = nn.Linear(img_dim, num_classes)
self.classifier_cd = nn.Linear(meta_dim, num_classes)
# Fused classifier head
self.classifier_fused = HTClassifier(fusion_dim, num_classes, dropout)
def reset_se_stats(self):
# ------------------------------------------------------------------
# SE helpers
# ------------------------------------------------------------------
def reset_se_stats(self) -> None:
"""Call at epoch start."""
if getattr(self, "se_log", None):
self.se_log.reset()
@@ -55,41 +101,205 @@ class Bridge(nn.Module):
return self.se_log.get(reset=reset)
return None
def _compute_fused(self, img_feats, md_feats):
"""Return z_fused embedding (before classifier_fused). Used by encode() and forward()."""
hi = self.ln_img(self.W_img(img_feats))
hm = self.ln_md(self.W_md(md_feats))
fused = hi * hm
# ------------------------------------------------------------------
# Core fusion
# ------------------------------------------------------------------
def _compute_fused(self, embeddings: list[torch.Tensor]) -> torch.Tensor:
"""Return z_fused embedding (before classifier_fused)."""
assert len(embeddings) == len(self.W), (
f"Bridge expects {len(self.W)} embeddings, got {len(embeddings)}"
)
h = self.ln[0](self.W[0](embeddings[0]))
for i in range(1, len(embeddings)):
h = h * self.ln[i](self.W[i](embeddings[i]))
if self.se is not None:
fused, gates = self.se(fused)
h, gates = self.se(h)
if self.se_log.enabled:
self.se_log.accumulate(gates)
return fused
return h
def encode(self, img_feats, md_feats) -> torch.Tensor:
"""Return z_fused embedding without applying the classifier head."""
assert self.mode == "fused", "encode() only valid in fused mode"
return self._compute_fused(img_feats, md_feats)
# ------------------------------------------------------------------
# N-tower API
# ------------------------------------------------------------------
def forward(self, img_feats, md_feats):
out_img = None if self.mode == "clinical_only" else self.classifier_img(img_feats)
out_md = None if self.mode == "image_only" else self.classifier_cd(md_feats)
def fuse(
self, embeddings: list[torch.Tensor]
) -> tuple[torch.Tensor, list[torch.Tensor]]:
"""
N-tower forward pass.
if self.mode == "fused":
fused = self._compute_fused(img_feats, md_feats)
out_f = self.classifier_fused(fused)
return out_f, out_img, out_md
# if ablation modes:
if self.mode == "image_only":
return out_img, out_img, None
if self.mode == "clinical_only":
return out_md, None, out_md
Parameters
----------
embeddings : list of Tensor — one per tower slot (same order as tower_dims).
Returns
-------
logits_fused : Tensor [B, num_classes]
aux_logits : list of Tensor — one per tower slot, each [B, num_classes]
"""
z_fused = self._compute_fused(embeddings)
logits_fused = self.classifier_fused(z_fused)
aux = [head(e) for head, e in zip(self.aux_heads, embeddings)]
return logits_fused, aux
def encode(self, embeddings: list[torch.Tensor]) -> torch.Tensor:
"""Return z_fused without applying the classifier head."""
return self._compute_fused(embeddings)
def set_phase(self, phase: str) -> None:
"""
Set requires_grad on bridge sub-modules according to training phase.
- ``cd_warmup`` — freeze everything in the bridge
- ``tower_warmup`` — aux_heads trainable, projections + fused head frozen
- ``fused_warmup`` — projections + fused head trainable, aux_heads frozen
- ``main`` / other — everything trainable
"""
def _rg(module, enabled):
for p in module.parameters():
p.requires_grad = enabled
if phase == "cd_warmup":
_rg(self, False)
return
if phase == "tower_warmup":
for head in self.aux_heads:
_rg(head, True)
for W_i in self.W:
_rg(W_i, False)
for ln_i in self.ln:
_rg(ln_i, False)
_rg(self.classifier_fused, False)
if self.se is not None:
_rg(self.se, False)
return
if phase == "fused_warmup":
for head in self.aux_heads:
_rg(head, False)
for W_i in self.W:
_rg(W_i, True)
for ln_i in self.ln:
_rg(ln_i, True)
_rg(self.classifier_fused, True)
if self.se is not None:
_rg(self.se, True)
return
_rg(self, True)
def forward(
self, embeddings: list[torch.Tensor]
) -> tuple[torch.Tensor, list[torch.Tensor]]:
"""N-tower forward. Delegates to fuse()."""
return self.fuse(embeddings)
# ---------------------------------------------------------------------------
# HyperBridge — higher-order bridge over HT module outputs
# ---------------------------------------------------------------------------
class HyperBridge(nn.Module):
"""Higher-order bridge that fuses z_fused embeddings from multiple HT modules.
Operates at the HT output level (z_fused from each HT's Bridge.encode())
rather than raw tower embedding level.
Modes
-----
embedding_mlp (default)
Concatenate all z_fused inputs → MLP → logits.
Analogous to EmbeddingMLPEnsembleHT, generalised to N inputs.
``Linear(N*fusion_dim → hidden_dim) → ReLU → Dropout → Linear(hidden_dim → num_classes)``
classic_bridge
Project each input to ``hidden_dim``, Hadamard product, HTClassifier.
Analogous to Bridge operating at the HT level — handles inputs of
differing dims via per-input projection layers.
``W[i](z_i) → LayerNorm → Hadamard → ReLU → Dropout → Linear(hidden_dim → num_classes)``
Both modes expose per-input auxiliary HTClassifier heads for BCD-style training.
Parameters
----------
input_dims : ordered dict {name: dim} for each HT input.
In embedding_mlp mode, dims may differ.
In classic_bridge mode, all dims must be equal (shared space).
num_classes : output classes
hidden_dim : hidden dim for the embedding_mlp MLP head
mode : "embedding_mlp" | "classic_bridge"
dropout : dropout throughout
"""
def __init__(
self,
input_dims: dict[str, int],
num_classes: int,
hidden_dim: int = 256,
mode: str = "embedding_mlp",
dropout: float = 0.3,
):
super().__init__()
self.input_names = list(input_dims.keys())
self.mode = mode
dims = list(input_dims.values())
if mode == "embedding_mlp":
total_dim = sum(dims)
self.head = nn.Sequential(
nn.Linear(total_dim, hidden_dim), nn.ReLU(),
nn.Dropout(dropout), nn.Linear(hidden_dim, num_classes),
)
elif mode == "classic_bridge":
# Project each input to shared fusion_dim space, then Hadamard
self.W = nn.ModuleList([nn.Linear(d, hidden_dim) for d in dims])
self.ln = nn.ModuleList([nn.LayerNorm(hidden_dim) for _ in dims])
self.head = HTClassifier(hidden_dim, num_classes, dropout)
else:
raise ValueError(f"Unknown HyperBridge mode: {mode!r}")
# Per-input auxiliary classifiers (both modes)
self.aux_heads = nn.ModuleList([
HTClassifier(d, num_classes, dropout) for d in dims
])
def forward(
self,
inputs: dict[str, torch.Tensor],
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
"""Fuse HT-level embeddings.
Parameters
----------
inputs : {name: z_fused [B, dim]} — z_fused from each HT's encode()
Returns
-------
logits : [B, num_classes]
aux_dict : {name: [B, num_classes]} — per-input aux head logits
"""
ordered = [inputs[name] for name in self.input_names]
if self.mode == "embedding_mlp":
logits = self.head(torch.cat(ordered, dim=1))
else: # classic_bridge: project → Hadamard → classify
h = self.ln[0](self.W[0](ordered[0]))
for i in range(1, len(ordered)):
h = h * self.ln[i](self.W[i](ordered[i]))
logits = self.head(h)
aux = {name: head(z) for name, head, z
in zip(self.input_names, self.aux_heads, ordered)}
return logits, aux
# ---------------------------------------------------------------------------
# VoteBridge
# ---------------------------------------------------------------------------
class VoteBridge(nn.Module):
def __init__(self, num_classes):
super().__init__()
self.vote_combiner = nn.Linear(num_classes * 2, num_classes) # two sets of logits
self.vote_combiner = nn.Linear(num_classes * 2, num_classes)
def forward(self, out_img, out_md):
votes = torch.cat([out_img, out_md], dim=1)
+162
View File
@@ -0,0 +1,162 @@
"""clinical_towers — ClinicalEncoder and ClinicalDataTower.
Self-contained: defines ClinicalEncoder directly (does not import it from
towers.py). Imports only TowerBase from towerbase plus infrastructure
(SEBlock, DataBundle).
"""
from __future__ import annotations
from typing import Optional
import torch
from torch import nn
from v3.classes.towerbase import TowerBase
from v3.classes.SE_attention import SEBlock
from v3.classes.data_bundle import DataBundle
# ---------------------------------------------------------------------------
# ClinicalEncoder — MLP over tabular features
# ---------------------------------------------------------------------------
class ClinicalEncoder(nn.Module):
"""MLP over DataBundle.vectorize_row outputs (converts to torch inside tower)."""
def __init__(
self,
clinical_data: DataBundle,
hidden_dim: int = 128,
dropout: float = 0.1,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
):
super().__init__()
self.feature_dim = clinical_data.feature_dim
self.out_dim = hidden_dim
# Two-block MLP so we can optionally freeze/thaw per block.
self.block0 = nn.Sequential(
nn.Linear(self.feature_dim, hidden_dim),
nn.LayerNorm(hidden_dim),
nn.ReLU(inplace=True),
nn.Dropout(dropout),
)
self.block1 = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(inplace=True),
)
self.net = nn.Sequential(self.block0, self.block1)
self.tower_ln = nn.LayerNorm(hidden_dim) if se_pre_norm else nn.Identity()
self.tower_se = (
SEBlock(hidden_dim, reduction=se_reduction, residual=True)
if use_se
else None
)
def forward(self, meta_np_or_torch) -> torch.Tensor:
if isinstance(meta_np_or_torch, torch.Tensor):
x = meta_np_or_torch
else:
x = torch.as_tensor(meta_np_or_torch, dtype=torch.float32)
h = self.net(x)
if self.tower_se is not None:
h, _ = self.tower_se(self.tower_ln(h))
return h
def set_freeze_ratio(self, ratio: float):
"""Optionally freeze earliest blocks of the MLP."""
r = max(0.0, min(1.0, float(ratio)))
for p in self.block0.parameters():
p.requires_grad = True
for p in self.block1.parameters():
p.requires_grad = True
if r >= 0.5:
for p in self.block0.parameters():
p.requires_grad = False
if r >= 1.0:
for p in self.block1.parameters():
p.requires_grad = False
# ---------------------------------------------------------------------------
# ClinicalDataTower — TowerBase implementation
# ---------------------------------------------------------------------------
class ClinicalDataTower(TowerBase, nn.Module):
"""
TowerBase implementation for the clinical metadata modality.
Wraps ClinicalEncoder (MLP over tabular features).
Contributes one embedding per eye slot: [z_cd].
Implements ``cd_warmup_embedding`` so train_towers_epoch can identify
this tower for cd_warmup phase via duck typing rather than isinstance checks.
"""
def __init__(
self,
*,
clinical_data,
cd_hidden_dim: int = 128,
cd_dropout: float = 0.1,
use_se: bool = False,
):
nn.Module.__init__(self)
self._encoder = ClinicalEncoder(
clinical_data=clinical_data,
hidden_dim=cd_hidden_dim,
dropout=cd_dropout,
use_se=use_se,
)
@property
def out_dim(self) -> int:
return self._encoder.out_dim
@property
def embed_dims(self) -> list[int]:
return [self._encoder.out_dim]
def set_phase(self, phase: str) -> None:
enabled = phase not in ("fused_warmup",)
for p in self._encoder.parameters():
p.requires_grad = enabled
def cd_warmup_embedding(
self,
batch: dict,
*,
device: torch.device,
) -> Optional[torch.Tensor]:
"""Return clinical embedding for slot 1, or None if matrix_1 is absent."""
m = batch.get("matrix_1")
if not torch.is_tensor(m):
return None
return self._encoder(m.to(device))
def embed_batch(
self,
batch: dict,
*,
device: torch.device,
slot: int = 1,
) -> list[torch.Tensor]:
m = batch.get(f"matrix_{slot}")
if not torch.is_tensor(m):
raise ValueError(f"ClinicalDataTower.embed_batch: matrix_{slot} is missing or not a tensor")
return [self._encoder(m.to(device))]
def prepare_fold(
self,
*,
eye_train,
bilat_train,
bilat_val,
bilat_test,
image_preprocessor,
image_cache,
device,
args,
) -> None:
pass # shares loader with ImageTower; no per-fold setup needed
-276
View File
@@ -1,276 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Dict, Iterable, List, Optional
import json
from v3.classes.papila_data import PapilaData
@dataclass
class ImportSpec:
id: str
class_name: str
params: Dict[str, Any]
@dataclass
class DataSourceSpec:
node_id: str
label: str
output_type: str
source: Optional[Dict[str, Any]]
source_ref: Optional[Dict[str, Any]]
@dataclass
class TransformSpec:
node_id: str
label: str
transform_type: str
params: Dict[str, Any]
@dataclass
class LoaderSpec:
node_id: str
label: str
input_type: str
input_index: str
input_key: str
output_key: str
transforms: List[TransformSpec]
data_source: Optional[DataSourceSpec]
@dataclass
class TowerSpec:
node_id: str
label: str
tower_type: str
params: Dict[str, Any]
@dataclass
class BridgeSpec:
node_id: str
label: str
method: str
params: Dict[str, Any]
@dataclass
class ClassifierSpec:
node_id: str
label: str
@dataclass
class ConfigAssembly:
raw: Dict[str, Any]
imports: Dict[str, ImportSpec]
data_sources: Dict[str, DataSourceSpec]
transforms: Dict[str, TransformSpec]
loaders: Dict[str, LoaderSpec]
towers: Dict[str, TowerSpec]
bridges: Dict[str, BridgeSpec]
classifiers: Dict[str, ClassifierSpec]
def load_config(path: Path) -> Dict[str, Any]:
payload = json.loads(Path(path).read_text())
if not isinstance(payload, dict):
raise ValueError("Config JSON must be an object.")
return payload
def assemble_config(path: Path) -> ConfigAssembly:
config = load_config(path)
meta = config.get("meta", {})
imports = _build_imports(meta.get("imports", []))
nodes = {node["id"]: node for node in config.get("nodes", [])}
edges = config.get("edges", [])
data_sources: Dict[str, DataSourceSpec] = {}
transforms: Dict[str, TransformSpec] = {}
loaders: Dict[str, LoaderSpec] = {}
towers: Dict[str, TowerSpec] = {}
bridges: Dict[str, BridgeSpec] = {}
classifiers: Dict[str, ClassifierSpec] = {}
for node in nodes.values():
ntype = node.get("type")
if ntype == "data":
data_sources[node["id"]] = DataSourceSpec(
node_id=node["id"],
label=node.get("label", ""),
output_type=node.get("outputType", ""),
source=node.get("source"),
source_ref=node.get("sourceRef"),
)
elif ntype == "transform":
transforms[node["id"]] = TransformSpec(
node_id=node["id"],
label=node.get("label", ""),
transform_type=node.get("transformType", ""),
params=_extract_transform_params(node),
)
elif ntype == "loader":
loaders[node["id"]] = LoaderSpec(
node_id=node["id"],
label=node.get("label", ""),
input_type=node.get("inputType", ""),
input_index=node.get("inputIndex", ""),
input_key=node.get("inputKey", ""),
output_key=node.get("outputKey", ""),
transforms=[],
data_source=None,
)
elif ntype in ("image_tower", "metadata_tower"):
towers[node["id"]] = TowerSpec(
node_id=node["id"],
label=node.get("label", ""),
tower_type=node.get("towerType", "image" if ntype == "image_tower" else "clinical data"),
params=_extract_tower_params(node),
)
elif ntype == "bridge":
bridges[node["id"]] = BridgeSpec(
node_id=node["id"],
label=node.get("label", ""),
method=node.get("bridgeMethod", "fusion"),
params=_extract_bridge_params(node),
)
elif ntype == "classifier":
classifiers[node["id"]] = ClassifierSpec(
node_id=node["id"],
label=node.get("label", ""),
)
# attach transforms + data sources to loaders by walking upstream
for loader_id, loader in loaders.items():
chain = _upstream_chain(loader_id, nodes, edges)
for node_id in reversed(chain):
if node_id in transforms:
loader.transforms.append(transforms[node_id])
if node_id in data_sources:
loader.data_source = data_sources[node_id]
return ConfigAssembly(
raw=config,
imports=imports,
data_sources=data_sources,
transforms=transforms,
loaders=loaders,
towers=towers,
bridges=bridges,
classifiers=classifiers,
)
def resolve_imports(assembly: ConfigAssembly) -> Dict[str, Any]:
resolved: Dict[str, Any] = {}
for import_id, spec in assembly.imports.items():
if spec.class_name == "PapilaData":
params = spec.params
resolved[import_id] = PapilaData.from_dirs(
image_dir=params.get("image_dir", "Papila/FundusImages"),
clinical_dir=params.get("clinical_dir", "Papila/ClinicalData"),
label_col=params.get("label_col", "Diagnosis"),
cat_cols=params.get("cat_cols", ["Gender", "Phakic/Pseudophakic"]),
)
else:
raise ValueError(f"Unsupported import class {spec.class_name!r}")
return resolved
def _build_imports(entries: Iterable[Dict[str, Any]]) -> Dict[str, ImportSpec]:
specs: Dict[str, ImportSpec] = {}
for entry in entries or []:
import_id = entry.get("id")
if not import_id:
continue
specs[import_id] = ImportSpec(
id=import_id,
class_name=entry.get("className", ""),
params=entry.get("params", {}) or {},
)
return specs
def _extract_transform_params(node: Dict[str, Any]) -> Dict[str, Any]:
return {
"transformType": node.get("transformType"),
"roiMaskSource": node.get("roiMaskSource"),
"roiScale": node.get("roiScale"),
"roiTargetSize": node.get("roiTargetSize"),
"roiFallback": node.get("roiFallback"),
"centerCropSize": node.get("centerCropSize"),
"jitterHFlip": node.get("jitterHFlip"),
"jitterVFlip": node.get("jitterVFlip"),
"jitterRotation": node.get("jitterRotation"),
"jitterColorEnabled": node.get("jitterColorEnabled"),
"jitterColor": node.get("jitterColor"),
"resizeSize": node.get("resizeSize"),
}
def _extract_tower_params(node: Dict[str, Any]) -> Dict[str, Any]:
if node.get("towerType") == "clinical data":
return {
"hidden_dim": node.get("mdHiddenDim"),
"dropout": node.get("mdDropout"),
"use_se": node.get("mdUseSe"),
"se_reduction": node.get("mdSeReduction"),
"se_pre_norm": node.get("mdSePreNorm"),
"freeze_ratio": node.get("mdFreezeRatio"),
}
return {
"backbone": node.get("imageBackbone"),
"freeze_ratio": node.get("imageFreezeRatio"),
"augment": node.get("imageAugment"),
"geometry_dim": node.get("imageGeometryDim"),
"use_se": node.get("imageUseSe"),
"se_reduction": node.get("imageSeReduction"),
"se_pre_norm": node.get("imageSePreNorm"),
}
def _extract_bridge_params(node: Dict[str, Any]) -> Dict[str, Any]:
return {
"fusion_dim": node.get("bridgeFusionDim"),
"use_se": node.get("bridgeUseSe"),
"se_reduction": node.get("bridgeSeReduction"),
"se_pre_norm": node.get("bridgeSePreNorm"),
}
def _edge_from(edge: Dict[str, Any]) -> Optional[str]:
return edge.get("from") or edge.get("source")
def _edge_to(edge: Dict[str, Any]) -> Optional[str]:
return edge.get("to") or edge.get("target")
def _upstream_chain(start_id: str, nodes: Dict[str, Dict[str, Any]], edges: List[Dict[str, Any]]) -> List[str]:
chain: List[str] = []
visited = set()
current = start_id
while True:
if current in visited:
break
visited.add(current)
incoming = [edge for edge in edges if _edge_to(edge) == current]
if not incoming:
break
# prefer first incoming edge for now
current = _edge_from(incoming[0])
if not current:
break
chain.append(current)
node = nodes.get(current)
if node and node.get("type") == "data":
break
return chain
-115
View File
@@ -1,115 +0,0 @@
from torch.utils.data import Dataset
from PIL import Image
import numpy as np
import torch
class ClinicalDataset(Dataset):
"""Generic dataset wrapping a DataBundle-like instance.
Returns (img_tensor, meta_tensor, label)."""
def __init__(
self,
clinical_data,
img_transform,
meta_transform=None,
image_preprocessor=None,
geometry_provider=None,
geometry_dim: int = 0,
image_cache: "dict | None" = None,
):
self.clinical = clinical_data
self.transform_image = img_transform
self.meta_transform = meta_transform or (lambda x: x)
self.image_preprocessor = image_preprocessor
self.geometry_provider = geometry_provider
self.geometry_dim = geometry_dim if geometry_provider is not None else 0
self.image_cache = image_cache
def __len__(self):
return len(self.clinical.df)
def __getitem__(self, idx: int):
row = self.clinical.df.iloc[idx]
# load & transform image
img_path = self.clinical.get_image_path(row)
cache_key = str(img_path)
if self.image_cache is not None and cache_key in self.image_cache:
orig_img = Image.fromarray(self.image_cache[cache_key])
else:
orig_img = Image.open(img_path).convert("RGB")
if self.image_cache is not None:
self.image_cache[cache_key] = np.asarray(orig_img, dtype=np.uint8)
img = orig_img
if self.image_preprocessor is not None:
img = self.image_preprocessor(img, img_path)
img_t = self.transform_image(img)
# encode & transform metadata
meta = self.clinical.encode_metadata(row)
meta_t = self.meta_transform(meta)
# label
label = self.clinical.get_label(row)
if self.geometry_dim > 0:
features = None
if self.geometry_provider is not None and hasattr(self.geometry_provider, "geometry_features"):
features = self.geometry_provider.geometry_features(orig_img, img_path)
if features is None:
geom_vec = torch.zeros(self.geometry_dim, dtype=torch.float32)
else:
features = np.asarray(features, dtype=np.float32)
if features.shape[0] != self.geometry_dim:
geom_vec = torch.zeros(self.geometry_dim, dtype=torch.float32)
else:
geom_vec = torch.from_numpy(features)
return img_t, meta_t, geom_vec, label
return img_t, meta_t, label
# ---------------------------------------------------------------------------
# _ClinicalView — shim used by V2HyperTower._run_fold
# ---------------------------------------------------------------------------
from .data_bundle import DataBundle # noqa: E402
class _ClinicalView:
"""Minimal shim so ClinicalDataset can iterate an epoch-specific DataFrame
while still delegating encoding/paths/labels to the DataBundle object."""
def __init__(self, base: DataBundle, df):
self.base = base
self.df = df
@property
def image_dir(self):
return self.base.image_dir
@property
def clinical_dir(self):
return self.base.clinical_dir
@property
def id_cols(self):
return ("Patient ID", "eyeID")
@property
def label_col(self):
return self.base.label_col
@property
def filename_template(self):
return getattr(self.base, "filename_template", "RET{pid:03d}{eye}.jpg")
@property
def dim(self):
return self.base.feature_dim
def encode_metadata(self, row):
vec = self.base.vectorize_row(row)
return torch.as_tensor(vec, dtype=torch.float32)
def get_image_path(self, row):
return self.base.get_image_path(row)
def get_label(self, row):
return int(row[self.base.label_col])
-119
View File
@@ -1,119 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Iterable, List, Sequence, Tuple, Union
import re
import pandas as pd
@dataclass
class RegexFilter:
pattern: str
flags: int = 0
def apply_paths(self, paths: Sequence[str]) -> Tuple[List[str], List[str]]:
if not self.pattern:
return list(paths), []
try:
regex = re.compile(self.pattern, self.flags)
except re.error as err:
return list(paths), [f'Invalid regex "{self.pattern}": {err}']
filtered = [p for p in paths if regex.search(p)]
return filtered, []
@dataclass
class ColumnFilter:
column: str
operator: str
value: str
case_insensitive: bool = True
def apply_df(self, df: pd.DataFrame) -> Tuple[pd.DataFrame, List[str]]:
warnings: List[str] = []
if not self.column:
return df, ["Column filter missing column name."]
columns = list(df.columns)
col_index = _resolve_column_index(columns, self.column, warnings)
if col_index is None:
return df, warnings
col_name = columns[col_index]
if self.value is None or self.value == "":
return df, [f'Column filter "{self.column}" missing value.']
series = df[col_name]
mask = series.apply(
lambda cell: compare_cell(
cell, self.value, self.operator, case_insensitive=self.case_insensitive
)
)
return df[mask], warnings
FilterSpec = Union[RegexFilter, ColumnFilter]
def apply_regex_filters(paths: Sequence[str], filters: Iterable[RegexFilter]) -> Tuple[List[str], List[str]]:
filtered = list(paths)
warnings: List[str] = []
for filt in filters:
filtered, warn = filt.apply_paths(filtered)
warnings.extend(warn)
return filtered, warnings
def apply_column_filters(df: pd.DataFrame, filters: Iterable[ColumnFilter]) -> Tuple[pd.DataFrame, List[str]]:
filtered = df
warnings: List[str] = []
for filt in filters:
filtered, warn = filt.apply_df(filtered)
warnings.extend(warn)
return filtered, warnings
def compare_cell(cell, raw_value: str, operator: str, case_insensitive: bool = True) -> bool:
cell_str = "" if cell is None else str(cell).strip()
value_str = "" if raw_value is None else str(raw_value).strip()
if case_insensitive:
cell_str = cell_str.lower()
value_str = value_str.lower()
if operator == "=":
return cell_str == value_str
if operator == "!=":
return cell_str != value_str
cell_num = _to_float(cell_str)
value_num = _to_float(value_str)
if cell_num is None or value_num is None:
return False
if operator == ">":
return cell_num > value_num
if operator == ">=":
return cell_num >= value_num
if operator == "<":
return cell_num < value_num
if operator == "<=":
return cell_num <= value_num
return False
def _resolve_column_index(columns: Sequence[str], column: str, warnings: List[str]) -> int | None:
try:
return columns.index(column)
except ValueError:
lower = column.lower()
matches = [idx for idx, col in enumerate(columns) if str(col).lower() == lower]
if matches:
if len(matches) > 1:
warnings.append(
f'Column "{column}" matched multiple headers; using "{columns[matches[0]]}".'
)
return matches[0]
warnings.append(f'Column "{column}" not found.')
return None
def _to_float(value: str) -> float | None:
try:
return float(value)
except (TypeError, ValueError):
return None
@@ -1,19 +1,22 @@
"""Segmentation-map CNN for glaucoma grading.
"""geometry_towers — GeometryTower and all segmentation-map infrastructure.
Trains a CNN on combined disc/cup segmentation maps pixel values
0 = background, 1 = disc (rim only), 2 = cup
instead of raw RGB fundus images, forcing the model to learn purely
from optic nerve head geometry (CDR, rim width, cup location, etc.).
Self-contained: absorbs everything that was in seg_cnn.py so that file can
eventually be removed. Does not import from seg_cnn.py or any other tower file.
Imports only TowerBase from towerbase plus standard infrastructure.
Two segmentation sources are supported:
gt rasterise expert contour/mask annotations directly (pure NumPy/PIL,
no CUDA safe in DataLoader worker processes)
unet run a trained UNetSegmenter on the raw fundus image
Usage (import from training script):
from v3.classes.seg_cnn import SegMapRecord, SegMapDataset, SegCNN, seg_map_to_tensor
Contents
--------
SegMapRecord labelled-eye data record
_combine_masks merge disc/cup binary masks 3-class label map
crop_to_disc tight bounding-box crop
seg_map_to_tensor (H,W) uint8 (C,H,W) float32 tensor
load_gt_masks load GT disc+cup masks from contour/mask files
UNetFineTuneDataset Dataset for fine-tuning the UNet on GT annotations
precompute_unet_seg_maps batch UNet inference helper
SegMapDataset Dataset yielding (seg_tensor, label) pairs
SegCNN pretrained CNN adapted for segmentation-map input
GeometryTower TowerBase implementation (the main class to use)
"""
from __future__ import annotations
from dataclasses import dataclass
@@ -29,6 +32,10 @@ from torch.utils.data import Dataset
from torchvision import models, transforms
from tqdm import tqdm
import pandas as pd
from v3.classes.towerbase import TowerBase
# ---------------------------------------------------------------------------
# Data record
@@ -53,8 +60,7 @@ class SegMapRecord:
# ---------------------------------------------------------------------------
def _combine_masks(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray:
"""
Combine binary disc and cup masks into a 3-class label map.
"""Combine binary disc and cup masks into a 3-class label map.
Returns a uint8 array with values:
0 background
@@ -69,8 +75,7 @@ def _combine_masks(disc_mask: np.ndarray, cup_mask: np.ndarray) -> np.ndarray:
def crop_to_disc(seg_map: np.ndarray) -> np.ndarray:
"""
Crop a seg map tightly to the disc bounding box.
"""Crop a seg map tightly to the disc bounding box.
The disc is anywhere seg_map > 0 (i.e. rim or cup).
Returns the original array unchanged if no disc is found.
@@ -89,8 +94,7 @@ def seg_map_to_tensor(
channels: int,
target_size: int,
) -> torch.Tensor:
"""
Convert an (H, W) seg map with values {0, 1, 2} to a float tensor.
"""Convert an (H, W) seg map with values {0, 1, 2} to a float tensor.
channels=1 (1, H, W) float in [0, 1] (values 0/0.5/1.0)
channels=3 (3, H, W) one-hot binary channels [bg, disc_rim, cup]
@@ -118,13 +122,15 @@ def seg_map_to_tensor(
def _load_contour(path: Path) -> np.ndarray:
"""Load x,y contour pairs from a whitespace- or comma-delimited text file."""
arr = np.zeros((0, 2), dtype=np.float32)
for delimiter in (",", None):
try:
arr = np.loadtxt(str(path), delimiter=delimiter, comments="#", dtype=np.float32)
if arr.size > 0:
candidate = np.loadtxt(str(path), delimiter=delimiter, comments="#", dtype=np.float32)
if candidate.size > 0:
arr = candidate
break
except Exception:
arr = np.zeros((0, 2), dtype=np.float32)
pass
if arr.size == 0 or arr.ndim == 1:
return np.zeros((0, 2), dtype=np.float32)
if arr.shape[1] < 2:
@@ -135,13 +141,10 @@ def _load_contour(path: Path) -> np.ndarray:
def _contour_to_mask(
coords: np.ndarray, image_size: Tuple[int, int], target_size: int
) -> np.ndarray:
"""
Rasterise a polygon defined by (x, y) coords into a binary mask.
"""Rasterise a polygon defined by (x, y) coords into a binary mask.
image_size is the (width, height) of the original fundus image the
coordinate space the contour was annotated in. The mask is drawn at
that resolution then resized to target_size, matching UNetSegmenter's
behaviour and avoiding off-canvas clipping.
coordinate space the contour was annotated in.
"""
if coords is None or len(coords) < 3:
return np.zeros((target_size, target_size), dtype=np.uint8)
@@ -155,8 +158,7 @@ def _contour_to_mask(
def _extract_masks_from_image(
mask_path: Path, target_size: int
) -> Tuple[np.ndarray, np.ndarray]:
"""
Extract disc and cup binary masks from a segmentation image file.
"""Extract disc and cup binary masks from a segmentation image file.
Handles both grayscale label images (e.g. REFUGE .bmp) and
RGB colour-coded masks. Returns (disc_mask, cup_mask) both at
@@ -166,7 +168,6 @@ def _extract_masks_from_image(
arr = np.array(raw)
if arr.ndim == 2:
# Grayscale: identify background from edge statistics
edges = np.concatenate([arr[0], arr[-1], arr[:, 0], arr[:, -1]])
bg_val = int(np.argmax(np.bincount(edges.astype(np.int64).clip(0, 255), minlength=256)))
disc_arr = (arr != bg_val).astype(np.uint8)
@@ -182,9 +183,7 @@ def _extract_masks_from_image(
img_rgb = raw.convert("RGB")
arr = np.array(img_rgb)
h, w, c = arr.shape
edges_rgb = np.concatenate(
[arr[0], arr[-1], arr[:, 0], arr[:, -1]], axis=0
)
edges_rgb = np.concatenate([arr[0], arr[-1], arr[:, 0], arr[:, -1]], axis=0)
edge_colors, edge_counts = np.unique(edges_rgb.reshape(-1, c), axis=0, return_counts=True)
bg_color = edge_colors[int(np.argmax(edge_counts))]
colors, counts = np.unique(arr.reshape(-1, c), axis=0, return_counts=True)
@@ -200,7 +199,6 @@ def _extract_masks_from_image(
cup_color = colors[order[1]]
cup_arr[np.all(arr == cup_color, axis=-1)] = 1
# Resize to target_size with nearest-neighbour to preserve binary values
def _resize(m: np.ndarray) -> np.ndarray:
pil = Image.fromarray((m > 0).astype(np.uint8) * 255)
pil = pil.resize((target_size, target_size), Resampling.NEAREST)
@@ -209,9 +207,8 @@ def _extract_masks_from_image(
return _resize(disc_arr), _resize(cup_arr)
def load_gt_masks(rec: "SegMapRecord", target_size: int) -> Tuple[np.ndarray, np.ndarray]:
"""
Load GT disc + cup masks for one record.
def load_gt_masks(rec: SegMapRecord, target_size: int) -> Tuple[np.ndarray, np.ndarray]:
"""Load GT disc + cup masks for one record.
Handles annotation_type "contour" (x,y text file) and "mask" (image file).
Returns (disc_mask, cup_mask) as uint8 arrays of shape (target_size, target_size).
@@ -219,7 +216,6 @@ def load_gt_masks(rec: "SegMapRecord", target_size: int) -> Tuple[np.ndarray, np
disc_mask: Optional[np.ndarray] = None
cup_mask: Optional[np.ndarray] = None
# Get original image size so contour coordinates are drawn in the right space
with Image.open(rec.image_path) as _img:
image_size = _img.size # (width, height)
@@ -246,7 +242,6 @@ def load_gt_masks(rec: "SegMapRecord", target_size: int) -> Tuple[np.ndarray, np
if cup_mask is None:
cup_mask = np.zeros((target_size, target_size), dtype=np.uint8)
# Structural prior: cup must lie within disc
cup_mask = (cup_mask > 0) & (disc_mask > 0)
return disc_mask.astype(np.uint8), cup_mask.astype(np.uint8)
@@ -256,11 +251,7 @@ def load_gt_masks(rec: "SegMapRecord", target_size: int) -> Tuple[np.ndarray, np
# ---------------------------------------------------------------------------
class UNetFineTuneDataset(Dataset):
"""
Loads (image_tensor, mask_tensor) pairs for fine-tuning the U-Net on
PAPILA GT annotations. Uses the same preprocessing as UNetSegmenter
so the fine-tuned weights are compatible with inference.
"""
"""Loads (image_tensor, mask_tensor) pairs for fine-tuning the U-Net."""
def __init__(
self,
@@ -292,7 +283,6 @@ class UNetFineTuneDataset(Dataset):
image = Image.open(rec.image_path).convert("RGB")
image = image.resize((self.target_size, self.target_size), Resampling.BILINEAR)
img_tensor = self._normalize(self.to_tensor(image))
disc_mask, cup_mask = load_gt_masks(rec, self.target_size)
mask_tensor = torch.from_numpy(
np.stack([disc_mask, cup_mask], axis=0).astype(np.float32)
@@ -301,20 +291,15 @@ class UNetFineTuneDataset(Dataset):
# ---------------------------------------------------------------------------
# U-Net precomputation (run once per full record list, not per fold)
# U-Net precomputation
# ---------------------------------------------------------------------------
def precompute_unet_seg_maps(
records: List["SegMapRecord"],
records: List[SegMapRecord],
segmenter,
threshold: float = 0.5,
) -> List[np.ndarray]:
"""
Run the U-Net on every record and return a list of combined seg maps.
Call this once before the CV loop and pass the results to each fold's
SegMapDataset via precomputed_seg_maps, so the U-Net isn't re-run per fold.
"""
"""Run the U-Net on every record and return a list of combined seg maps."""
to_tensor = transforms.ToTensor()
seg_maps = []
for rec in tqdm(records, desc="U-Net inference", unit="img", leave=False):
@@ -334,28 +319,11 @@ def precompute_unet_seg_maps(
# ---------------------------------------------------------------------------
# Dataset
# SegMapDataset
# ---------------------------------------------------------------------------
class SegMapDataset(Dataset):
"""
PyTorch Dataset that yields (seg_tensor, label) pairs.
Parameters
----------
records : list of SegMapRecord
target_size : CNN input spatial size (images are resized to this)
channels : 1 = single-channel label map; 3 = one-hot three channels
augment : apply random flips + rotation (for training set)
unet_segmenter : if provided, use U-Net predictions instead of GT masks;
must be a loaded UNetSegmenter with model weights set
unet_threshold : threshold for U-Net logit binary mask
seg_target_size: resolution at which GT masks are rasterised (or U-Net
output size). Default 512 matches UNetSegmenter default.
crop_to_disc : crop the seg map tightly to the disc bounding box before
resizing to target_size (default True eliminates the
background zeros that make up most of the full image)
"""
"""PyTorch Dataset that yields (seg_tensor, label) pairs."""
def __init__(
self,
@@ -366,7 +334,7 @@ class SegMapDataset(Dataset):
unet_segmenter=None,
unet_threshold: float = 0.5,
seg_target_size: int = 512,
crop_to_disc: bool = True,
crop_to_disc_flag: bool = True,
precomputed_seg_maps: Optional[List[np.ndarray]] = None,
) -> None:
self.records = records
@@ -374,7 +342,7 @@ class SegMapDataset(Dataset):
self.channels = channels
self.augment = augment
self.seg_target_size = seg_target_size
self.crop_to_disc = crop_to_disc
self.crop_to_disc_flag = crop_to_disc_flag
if precomputed_seg_maps is not None:
self._seg_maps = precomputed_seg_maps
@@ -385,13 +353,10 @@ class SegMapDataset(Dataset):
else:
self._seg_maps = None
# ------------------------------------------------------------------
def __len__(self) -> int:
return len(self.records)
# ------------------------------------------------------------------
def _augment(self, seg_map: np.ndarray) -> np.ndarray:
"""Random flips + 90° rotations (label-safe since NEAREST resize)."""
if np.random.rand() < 0.5:
seg_map = np.fliplr(seg_map)
if np.random.rand() < 0.5:
@@ -401,19 +366,16 @@ class SegMapDataset(Dataset):
seg_map = np.rot90(seg_map, k=k)
return np.ascontiguousarray(seg_map)
# ------------------------------------------------------------------
def __getitem__(self, idx: int):
rec = self.records[idx]
if self._seg_maps is not None:
seg_map = self._seg_maps[idx]
else:
disc_mask, cup_mask = load_gt_masks(rec, self.seg_target_size)
seg_map = _combine_masks(disc_mask, cup_mask)
if self.crop_to_disc:
if self.crop_to_disc_flag:
seg_map = crop_to_disc(seg_map)
if self.augment:
seg_map = self._augment(seg_map)
@@ -422,19 +384,24 @@ class SegMapDataset(Dataset):
# ---------------------------------------------------------------------------
# Model
# SegCNN
# ---------------------------------------------------------------------------
_SEGCNN_FEAT_DIM = {
"resnet18": 512,
"resnet50": 2048,
"efficientnet_b0": 1280,
}
class SegCNN(nn.Module):
"""
Pretrained CNN backbone adapted for segmentation-map input.
"""Pretrained CNN backbone adapted for segmentation-map input.
Parameters
----------
num_classes : output classes (2 for binary glaucoma grading)
backbone : "resnet18" | "resnet50" | "efficientnet_b0"
pretrained : initialise with ImageNet weights (recommended even for
non-RGB input transfer generalises across domains)
pretrained : initialise with ImageNet weights
in_channels : 1 (single label map) or 3 (one-hot channels)
dropout : dropout rate before the final classifier head
"""
@@ -448,7 +415,6 @@ class SegCNN(nn.Module):
dropout: float = 0.3,
) -> None:
super().__init__()
weights_arg = "DEFAULT" if pretrained else None
if backbone == "resnet18":
@@ -466,7 +432,6 @@ class SegCNN(nn.Module):
else:
raise ValueError(f"Unknown backbone: {backbone!r}")
# Adapt first conv layer if in_channels ≠ 3
if in_channels != 3:
first_conv = self._find_first_conv(base)
new_conv = nn.Conv2d(
@@ -478,7 +443,6 @@ class SegCNN(nn.Module):
bias=first_conv.bias is not None,
)
if pretrained:
# Average pretrained RGB weights across channel dim
with torch.no_grad():
new_conv.weight.copy_(
first_conv.weight.mean(dim=1, keepdim=True).expand_as(new_conv.weight)
@@ -491,7 +455,6 @@ class SegCNN(nn.Module):
nn.Linear(feat_dim, num_classes),
)
# ------------------------------------------------------------------
@staticmethod
def _find_first_conv(module: nn.Module) -> nn.Conv2d:
for m in module.modules():
@@ -501,7 +464,6 @@ class SegCNN(nn.Module):
@staticmethod
def _replace_first_conv(module: nn.Module, new_conv: nn.Conv2d) -> None:
"""Replace the first Conv2d in-place (handles resnet and efficientnet)."""
for name, child in module.named_children():
if isinstance(child, nn.Conv2d):
setattr(module, name, new_conv)
@@ -513,9 +475,304 @@ class SegCNN(nn.Module):
pass
raise RuntimeError("Could not replace first Conv2d")
# ------------------------------------------------------------------
def forward(self, x: torch.Tensor) -> torch.Tensor:
feats = self.backbone(x)
if feats.dim() > 2:
feats = feats.flatten(1)
return self.head(feats)
# ---------------------------------------------------------------------------
# GeometryTower — TowerBase implementation
# ---------------------------------------------------------------------------
class GeometryTower(TowerBase, nn.Module):
"""TowerBase implementation for the optic-disc/cup segmentation modality.
Encodes a 3-class disc/cup segmentation map (bg=0, rim=1, cup=2) through a
CNN backbone, contributing one spatial embedding to the bridge.
The seg map is produced from GT annotations (manifest-based) or from a
trained U-Net, depending on ``geometry_source``.
``prepare_fold`` builds a seg-map generator, pre-computes all maps for the
fold, and caches them keyed by image path. ``augment_samples`` then
injects ``seg_map_1`` / ``seg_map_2`` float32 numpy arrays (shape C×H×W)
into each sample dict so the DataLoader delivers them as tensors to
``embed_batch``.
Parameters
----------
backbone : CNN backbone "resnet18" | "resnet50" | "efficientnet_b0"
in_channels : 1 (label map) or 3 (one-hot disc/rim/cup channels)
pretrained : initialise backbone with ImageNet weights
frozen : if True, backbone is always frozen
target_size : spatial size the seg map tensor is resized to
seg_target_size : resolution at which GT masks are rasterised / U-Net runs
crop_to_disc : crop seg map tightly to disc bounding box before resizing
geometry_source : "gt" (manifest annotations) or "unet" (U-Net predictions)
manifest_path : path to the geometry manifest CSV (required)
weights_path : path to UNet checkpoint (required when source="unet")
unet_normalize : UNet normalisation mode (default "per_image")
unet_threshold : UNet mask threshold (default 0.5)
finetune_unet_epochs : epochs to fine-tune U-Net per fold (0 = disabled)
finetune_unet_lr : learning rate for U-Net fine-tuning
"""
def __init__(
self,
*,
backbone: str = "resnet18",
in_channels: int = 3,
pretrained: bool = True,
frozen: bool = False,
target_size: int = 224,
seg_target_size: int = 512,
crop_to_disc: bool = True,
geometry_source: str = "gt",
manifest_path=None,
weights_path=None,
unet_normalize: str = "per_image",
unet_threshold: float = 0.5,
finetune_unet_epochs: int = 0,
finetune_unet_lr: float = 1e-5,
):
nn.Module.__init__(self)
self._backbone_name = backbone
self._in_channels = in_channels
self._frozen = frozen
self._target_size = target_size
self._seg_target_size = seg_target_size
self._crop_to_disc = crop_to_disc
self._geometry_source = geometry_source
self._manifest_path = Path(manifest_path) if manifest_path is not None else None
self._weights_path = Path(weights_path) if weights_path is not None else None
self._unet_normalize = unet_normalize
self._unet_threshold = unet_threshold
self._finetune_unet_epochs = finetune_unet_epochs
self._finetune_unet_lr = finetune_unet_lr
self._out_dim = _SEGCNN_FEAT_DIM.get(backbone, 512)
self._seg_cnn = SegCNN(
num_classes=2,
backbone=backbone,
pretrained=pretrained,
in_channels=in_channels,
)
self._seg_cache: dict = {} # image_path_str → float32 (C, H, W) numpy array
# ------------------------------------------------------------------
# TowerBase interface
# ------------------------------------------------------------------
@property
def embed_dims(self) -> list[int]:
return [self._out_dim]
@property
def total_epochs(self) -> int:
return 0 if self._frozen else 1
def set_phase(self, phase: str) -> None:
trainable = not self._frozen and phase in ("tower_warmup", "main")
for p in self._seg_cnn.parameters():
p.requires_grad = trainable
def prepare_fold(
self,
*,
eye_train,
bilat_train,
bilat_val,
bilat_test,
image_preprocessor,
image_cache,
device,
args,
) -> None:
"""Build seg-map generator and pre-compute maps for all fold images."""
if self._manifest_path is None:
raise ValueError("GeometryTower requires manifest_path")
all_paths: dict = {}
for split in (eye_train, bilat_train, bilat_val, bilat_test):
for s in split:
for slot in ("image_1", "image_2"):
p = s.get(slot)
if p is not None:
all_paths[str(Path(p).resolve())] = None
if self._geometry_source == "unet":
self._prepare_fold_unet(list(all_paths.keys()), eye_train, device)
else:
self._prepare_fold_gt(list(all_paths.keys()))
def augment_samples(self, samples: list) -> list:
"""Inject ``seg_map_1`` / ``seg_map_2`` float32 arrays into each sample dict.
Arrays have shape (C, H, W) and are collated by the DataLoader into
(B, C, H, W) tensors delivered to ``embed_batch``.
"""
blank = np.zeros(
(self._in_channels, self._target_size, self._target_size), dtype=np.float32
)
for s in samples:
for img_slot, seg_slot in (("image_1", "seg_map_1"), ("image_2", "seg_map_2")):
img_path = s.get(img_slot)
if img_path is None:
continue
key = str(Path(img_path).resolve())
s[seg_slot] = self._seg_cache.get(key, blank)
return samples
def embed_batch(
self,
batch: dict,
*,
device: torch.device,
slot: int = 1,
) -> list[torch.Tensor]:
seg = batch.get(f"seg_map_{slot}")
if seg is None or not torch.is_tensor(seg):
ref = batch.get(f"image_{slot}")
bs = ref.shape[0] if torch.is_tensor(ref) else 1
return [torch.zeros(bs, self._out_dim, device=device)]
return [self._seg_cnn.backbone(seg.float().to(device))]
# ------------------------------------------------------------------
# Internal helpers
# ------------------------------------------------------------------
def _seg_map_to_array(self, seg_map: np.ndarray) -> np.ndarray:
"""Apply crop + resize and return a (C, H, W) float32 numpy array."""
if self._crop_to_disc:
seg_map = crop_to_disc(seg_map)
return seg_map_to_tensor(seg_map, self._in_channels, self._target_size).numpy()
def _prepare_fold_gt(self, image_paths: list) -> None:
"""Pre-compute GT seg maps from manifest annotations."""
manifest_df = pd.read_csv(self._manifest_path)
manifest_df["_img_key"] = manifest_df["image_path"].apply(
lambda p: str(Path(p).resolve())
)
manifest_index = manifest_df.set_index("_img_key").to_dict("index")
print(
f"[GeometryTower] pre-computing GT seg maps for {len(image_paths)} images...",
flush=True,
)
n_ok = 0
blank = np.zeros((self._seg_target_size, self._seg_target_size), dtype=np.uint8)
for img_path in image_paths:
entry = manifest_index.get(img_path)
if entry is None:
self._seg_cache[img_path] = self._seg_map_to_array(blank)
continue
rec = SegMapRecord(
sample_id="",
image_path=Path(img_path),
annotation_disc=Path(entry["annotation_disc"]),
annotation_cup=Path(entry["annotation_cup"]),
annotation_type_disc=entry["annotation_type_disc"],
annotation_type_cup=entry["annotation_type_cup"],
patient_id=0,
eye="",
label=0,
)
try:
disc_mask, cup_mask = load_gt_masks(rec, self._seg_target_size)
self._seg_cache[img_path] = self._seg_map_to_array(
_combine_masks(disc_mask, cup_mask)
)
n_ok += 1
except Exception:
self._seg_cache[img_path] = self._seg_map_to_array(blank)
print(f"[GeometryTower] {n_ok}/{len(image_paths)} GT seg maps computed", flush=True)
def _prepare_fold_unet(self, image_paths: list, eye_train: list, device) -> None:
"""Pre-compute U-Net seg maps, with optional per-fold fine-tuning."""
from v3.classes.unet_segmenter import UNetSegmenter
from torch.utils.data import DataLoader as _DL
if self._weights_path is None:
raise ValueError("GeometryTower(source='unet') requires weights_path")
segmenter = UNetSegmenter(
manifest_path=self._manifest_path,
normalize=self._unet_normalize,
)
state = torch.load(self._weights_path, map_location=segmenter.device)
segmenter.model.load_state_dict(state.get("model", state))
segmenter.model.to(segmenter.device).eval()
if self._finetune_unet_epochs > 0:
print(
f"[GeometryTower] fine-tuning U-Net for {self._finetune_unet_epochs} epochs...",
flush=True,
)
ft_loader = _DL(
UNetFineTuneDataset(
self._build_records_from_samples(eye_train),
target_size=segmenter.target_size,
normalize=self._unet_normalize,
),
batch_size=4, shuffle=True, num_workers=0,
)
optimizer = torch.optim.Adam(segmenter.model.parameters(), lr=self._finetune_unet_lr)
criterion = torch.nn.BCEWithLogitsLoss()
segmenter.model.train()
for _ in range(self._finetune_unet_epochs):
for images, masks in ft_loader:
images, masks = images.to(segmenter.device), masks.to(segmenter.device)
optimizer.zero_grad()
criterion(segmenter.model(images), masks).backward()
optimizer.step()
segmenter.model.eval()
print(
f"[GeometryTower] running U-Net inference on {len(image_paths)} images...",
flush=True,
)
records = [
SegMapRecord(
sample_id="", image_path=Path(p),
annotation_disc=Path(p), annotation_cup=Path(p),
annotation_type_disc="", annotation_type_cup="",
patient_id=0, eye="", label=0,
)
for p in image_paths
]
seg_maps = precompute_unet_seg_maps(records, segmenter, self._unet_threshold)
for img_path, seg_map in zip(image_paths, seg_maps):
self._seg_cache[img_path] = self._seg_map_to_array(seg_map)
print(f"[GeometryTower] {len(seg_maps)} U-Net seg maps cached", flush=True)
def _build_records_from_samples(self, samples: list) -> list:
"""Build SegMapRecord list from HyperTower sample dicts (for U-Net fine-tuning)."""
manifest_df = pd.read_csv(self._manifest_path)
manifest_df["_img_key"] = manifest_df["image_path"].apply(
lambda p: str(Path(p).resolve())
)
manifest_index = manifest_df.set_index("_img_key").to_dict("index")
records = []
for s in samples:
for slot in ("image_1", "image_2"):
p = s.get(slot)
if p is None:
continue
key = str(Path(p).resolve())
entry = manifest_index.get(key)
if entry is None:
continue
records.append(SegMapRecord(
sample_id="",
image_path=Path(p),
annotation_disc=Path(entry["annotation_disc"]),
annotation_cup=Path(entry["annotation_cup"]),
annotation_type_disc=entry["annotation_type_disc"],
annotation_type_cup=entry["annotation_type_cup"],
patient_id=int(s.get("patient_id", 0)),
eye=str(s.get("eye", "")),
label=int(s.get("label", 0)),
))
return records
File diff suppressed because it is too large Load Diff
+252
View File
@@ -0,0 +1,252 @@
"""image_towers — ImageEncoder, SiameseImageTower, and ImageTower (TowerBase).
Self-contained image-modality tower layer. No dependencies on other tower
files, bridge, or model classes. Clear contract: accepts an image batch,
returns a fixed-size embedding vector.
"""
from __future__ import annotations
import math
from typing import Optional
import torch
from torch import nn
from torch.utils.data import DataLoader
from v3.classes.towerbase import TowerBase, build_backbone
from v3.classes.backbones import BACKBONES
from v3.classes.SE_attention import SEBlock
# ---------------------------------------------------------------------------
# ImageEncoder — vision backbone → pooled feature vector
# ---------------------------------------------------------------------------
class ImageEncoder(nn.Module):
"""Vision backbone → pooled feature vector.
Wraps a torchvision backbone (default weights), strips the classifier,
and optionally appends an SE attention block and/or a geometry vector.
Parameters
----------
backbone : backbone key (see backbones.py)
freeze_ratio : fraction of early blocks to freeze in [0, 1]
use_se : apply SE attention over the pooled feature vector
augment : include random flip/rotation/jitter in the transform
geometry_dim : if > 0, concatenate a geometry vector of this length
"""
def __init__(
self,
backbone: str = "efficientnet_b0",
freeze_ratio: float = 0.0,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
augment: bool = True,
geometry_dim: int = 0,
):
super().__init__()
self.backbone, base_dim, self.transform = build_backbone(
backbone, freeze_ratio, augment=augment
)
self._name = backbone
key = (self._name or "").lower()
self._spec = BACKBONES[key]
self._blocks = self._spec.blocks(self.backbone)
self.base_dim = base_dim
self.geometry_dim = max(0, int(geometry_dim))
self.out_dim = self.base_dim + self.geometry_dim
self.tower_ln = nn.LayerNorm(self.base_dim) if se_pre_norm else nn.Identity()
self.tower_se = (
SEBlock(self.base_dim, reduction=se_reduction, residual=True)
if use_se else None
)
def forward(
self, x: torch.Tensor, geometry: Optional[torch.Tensor] = None
) -> torch.Tensor:
y = self.backbone(x)
assert y.dim() == 2 and y.size(1) == self.base_dim, (
f"Expected features [N,{self.base_dim}], got {tuple(y.shape)}"
)
if self.tower_se is not None:
y, _ = self.tower_se(self.tower_ln(y))
if self.geometry_dim > 0:
if geometry is None or geometry.numel() == 0:
geom = torch.zeros(y.size(0), self.geometry_dim, device=y.device, dtype=y.dtype)
else:
geom = geometry.unsqueeze(0) if geometry.dim() == 1 else geometry
geom = geom.to(device=y.device, dtype=y.dtype)
if geom.size(0) != y.size(0):
raise ValueError(f"Geometry batch size mismatch: {geom.size(0)} vs {y.size(0)}")
if geom.size(1) != self.geometry_dim:
raise ValueError(f"Expected geometry dim {self.geometry_dim}, got {geom.size(1)}")
y = torch.cat([y, geom], dim=1)
return y
def set_freeze_ratio(self, ratio: float) -> None:
"""Dynamically freeze earliest floor(N*ratio) backbone blocks."""
r = max(0.0, min(1.0, float(ratio)))
freeze_n = int(math.floor(len(self._blocks) * r))
for b in self._blocks:
for p in b.parameters():
p.requires_grad = True
for b in self._blocks[:freeze_n]:
for p in b.parameters():
p.requires_grad = False
# ---------------------------------------------------------------------------
# SiameseImageTower — shared-weight bilateral image encoder
# ---------------------------------------------------------------------------
class SiameseImageTower(nn.Module):
"""Shared-weight bilateral image tower.
Runs OD and OS images through a single shared backbone and returns
cat([f_mean, f_delta]) where:
f_mean = (f_od + f_os) / 2 — shared bilateral representation
f_delta = f_od - f_os — signed asymmetry (OD-relative)
out_dim = 2 × backbone_out_dim. When x_os is None the tower degrades
gracefully: f_mean = f_od, f_delta = zeros.
"""
def __init__(
self,
backbone: str = "efficientnet_b0",
freeze_ratio: float = 0.0,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
augment: bool = True,
):
super().__init__()
self._tower = ImageEncoder(
backbone=backbone,
freeze_ratio=freeze_ratio,
use_se=use_se,
se_reduction=se_reduction,
se_pre_norm=se_pre_norm,
augment=augment,
)
self.out_dim = self._tower.out_dim * 2
self.transform = self._tower.transform
def forward(
self,
x_od: torch.Tensor,
x_os: Optional[torch.Tensor] = None,
) -> torch.Tensor:
f_od = self._tower(x_od)
if x_os is None:
return torch.cat([f_od, torch.zeros_like(f_od)], dim=1)
f_os = self._tower(x_os)
return torch.cat([(f_od + f_os) * 0.5, f_od - f_os], dim=1)
def set_freeze_ratio(self, ratio: float) -> None:
self._tower.set_freeze_ratio(ratio)
# ---------------------------------------------------------------------------
# ImageTower — TowerBase implementation
# ---------------------------------------------------------------------------
class ImageTower(TowerBase, nn.Module):
"""TowerBase implementation for the fundus image modality.
Wraps ImageEncoder (backbone → pooled features).
Contributes one embedding per eye slot: [z_img].
"""
def __init__(
self,
*,
backbone: str,
freeze_ratio: float = 0.0,
augment: bool = True,
use_se: bool = False,
):
nn.Module.__init__(self)
self._encoder = ImageEncoder(backbone=backbone, freeze_ratio=freeze_ratio,
augment=augment, use_se=use_se)
self._train_loader = None
@property
def transform(self):
return self._encoder.transform
@property
def out_dim(self) -> int:
return self._encoder.out_dim
@property
def embed_dims(self) -> list[int]:
return [self._encoder.out_dim]
@property
def total_epochs(self) -> int:
return 0
@property
def train_loader(self) -> Optional[DataLoader]:
return self._train_loader
def set_phase(self, phase: str) -> None:
enabled = phase not in ("cd_warmup", "fused_warmup")
for p in self._encoder.parameters():
p.requires_grad = enabled
def embed_batch(
self,
batch: dict,
*,
device: torch.device,
slot: int = 1,
) -> list[torch.Tensor]:
x = batch.get(f"image_{slot}")
if not torch.is_tensor(x):
raise ValueError(f"ImageTower.embed_batch: image_{slot} missing or not a tensor")
return [self._encoder(x.to(device))]
def prepare_fold(
self,
*,
eye_train,
bilat_train,
bilat_val,
bilat_test,
image_preprocessor,
image_cache,
device,
args,
) -> None:
from v3.classes.loader_factory import (
build_balanced_sampler,
filter_eye_samples,
make_loader,
)
from v3.classes.profiles import build_papila_profile
profile_eye = build_papila_profile(
patient_col="Patient ID", label_col=args.label_col, sample_mode="eye"
)
slots_eye = profile_eye.slot_descriptors()
use_balanced = bool(getattr(args, "balanced_sampling", False))
_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,
)
eye_samples = filter_eye_samples(eye_train)
sampler = build_balanced_sampler(eye_samples) if use_balanced else None
self._train_loader = make_loader(
eye_samples, slots_eye,
image_transform=self.transform,
image_preprocessor=image_preprocessor,
shuffle=True, sampler=sampler, **loader_kw,
)
+2 -2
View File
@@ -4,7 +4,7 @@ from dataclasses import dataclass
from typing import Any, Callable, Optional
import torch
from torch.utils.data import DataLoader, WeightedRandomSampler
from torch.utils.data import DataLoader, Sampler, WeightedRandomSampler
from .network_manager import LoaderBundle, PatientSplit
from .slot_dataset import SlotDataset, slot_collate
@@ -210,7 +210,7 @@ def make_loader(
batch_size: int,
shuffle: bool,
num_workers: int,
sampler: Optional[WeightedRandomSampler] = None,
sampler: Optional[Sampler] = None,
persistent_workers: bool = False,
) -> DataLoader:
ds = SlotDataset(
-148
View File
@@ -1,148 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Callable, Optional
import torch
from torch import nn
from v3.classes.bridges import Bridge, VoteBridge
from v3.classes.towers import ImageTower, ClinicalTower
from .config_builder import ConfigAssembly
from .transforms import build_transform_chain
@dataclass
class V2ModelBundle:
image_tower: Optional[ImageTower]
metadata_tower: Optional[ClinicalTower]
bridge: Optional[nn.Module]
classifier: Optional[nn.Module]
image_transform: Optional[Callable]
matrix_transform: Optional[Callable]
def build_model_bundle(
assembly: ConfigAssembly,
clinical: Any,
*,
device: Optional[torch.device] = None,
strict: bool = True,
) -> V2ModelBundle:
"""
Build torch modules and input transforms from a V2 config assembly.
"""
image_tower_spec = _pick_tower(assembly, "image")
cd_tower_spec = _pick_tower(assembly, "clinical data")
bridge_spec = _pick_bridge(assembly)
image_loader = _pick_loader(assembly, input_type="image")
clinical_core = getattr(clinical, "clinical", clinical)
num_classes = _infer_num_classes(clinical)
img_tower = None
if image_tower_spec is not None:
img_tower = ImageTower(
backbone=image_tower_spec.params.get("backbone", "efficientnet_b0"),
freeze_ratio=float(image_tower_spec.params.get("freeze_ratio", 0.0) or 0.0),
use_se=bool(image_tower_spec.params.get("use_se", False)),
se_reduction=int(image_tower_spec.params.get("se_reduction", 16) or 16),
se_pre_norm=bool(image_tower_spec.params.get("se_pre_norm", True)),
augment=bool(image_tower_spec.params.get("augment", True)),
geometry_dim=int(image_tower_spec.params.get("geometry_dim", 0) or 0),
)
if device is not None:
img_tower = img_tower.to(device)
cd_tower = None
if cd_tower_spec is not None:
cd_tower = ClinicalTower(
clinical_core,
hidden_dim=int(cd_tower_spec.params.get("hidden_dim", 128) or 128),
dropout=float(cd_tower_spec.params.get("dropout", 0.1) or 0.1),
use_se=bool(cd_tower_spec.params.get("use_se", False)),
se_reduction=int(cd_tower_spec.params.get("se_reduction", 16) or 16),
se_pre_norm=bool(cd_tower_spec.params.get("se_pre_norm", True)),
)
if device is not None:
cd_tower = cd_tower.to(device)
bridge = None
if bridge_spec is not None and img_tower is not None and cd_tower is not None:
if bridge_spec.method == "consensus":
bridge = VoteBridge(num_classes=num_classes)
else:
bridge = Bridge(
img_dim=img_tower.out_dim,
meta_dim=cd_tower.out_dim,
num_classes=num_classes,
fusion_dim=int(bridge_spec.params.get("fusion_dim", 256) or 256),
mode="fused",
use_se=bool(bridge_spec.params.get("use_se", True)),
se_reduction=int(bridge_spec.params.get("se_reduction", 16) or 16),
se_pre_norm=bool(bridge_spec.params.get("se_pre_norm", True)),
)
if device is not None:
bridge = bridge.to(device)
classifier = None
if assembly.classifiers:
classifier = nn.Identity()
if device is not None:
classifier = classifier.to(device)
image_transform = None
if image_loader is not None and image_tower_spec is not None:
image_transform = build_transform_chain(
image_loader.transforms,
backbone_name=image_tower_spec.params.get("backbone", "efficientnet_b0"),
augment=bool(image_tower_spec.params.get("augment", True)),
strict=strict,
)
return V2ModelBundle(
image_tower=img_tower,
metadata_tower=cd_tower,
bridge=bridge,
classifier=classifier,
image_transform=image_transform,
matrix_transform=None,
)
def _pick_tower(assembly: ConfigAssembly, tower_type: str):
matches = [tower for tower in assembly.towers.values() if tower.tower_type == tower_type]
if not matches:
return None
if len(matches) > 1:
raise ValueError(f"Multiple {tower_type} towers found; only one is supported for now.")
return matches[0]
def _pick_bridge(assembly: ConfigAssembly):
if not assembly.bridges:
return None
if len(assembly.bridges) > 1:
raise ValueError("Multiple bridges found; only one is supported for now.")
return next(iter(assembly.bridges.values()))
def _pick_loader(assembly: ConfigAssembly, input_type: str):
matches = [loader for loader in assembly.loaders.values() if loader.input_type == input_type]
if not matches:
return None
if len(matches) > 1:
raise ValueError(f"Multiple loaders with input_type={input_type!r} found.")
return matches[0]
def _infer_num_classes(clinical: Any) -> int:
df = getattr(clinical, "df", None)
label_col = getattr(clinical, "label_col", None)
if df is None and hasattr(clinical, "clinical"):
df = clinical.clinical.df
label_col = clinical.clinical.label_col
if df is None or label_col is None or label_col not in df.columns:
return 2
return int(df[label_col].dropna().nunique())
-1074
View File
File diff suppressed because it is too large Load Diff
+455
View File
@@ -0,0 +1,455 @@
"""towerbase — unified TowerBase ABC, backbone factory, and modular training utilities.
This module is the structural backbone of the HyperTower v3 architecture.
It provides the abstract tower interface plus the training/eval helpers that
operate on any list of TowerBase instances.
Design principles
-----------------
* No concrete tower classes are defined here (ImageEncoder, ClinicalEncoder, etc.
live in their respective tower files).
* No imports from any tower file — this module is self-contained with respect to
the tower layer. Tower files import from here; this file does not import from them.
* Removing any tower file leaves this module fully intact.
* Duck-typing via optional TowerBase methods (e.g. cd_warmup_embedding) replaces
isinstance checks so new tower types never require changes here.
"""
from __future__ import annotations
import math
from abc import ABC, abstractmethod
from random import random
from dataclasses import dataclass, field
from typing import Optional
import numpy as np
import torch
import torch.nn.functional as F
from torch import nn
from torch.utils.data import DataLoader
from torchvision import transforms
from v3.classes.backbones import BACKBONES, list_names, load_backbone_weights
from v3.classes.bridges import Bridge
from random import random
# ---------------------------------------------------------------------------
# Cross-tower communication context
# ---------------------------------------------------------------------------
@dataclass
class EarlyPassContext:
eye_train: list[dict]
bilat_train: list[dict]
bilat_val: list[dict]
bilat_test: list[dict]
image_preprocessor: object
image_cache: object
device: torch.device
store: dict = field(default_factory=dict) # cross-tower key-value store
# ---------------------------------------------------------------------------
# Backbone factory
# ---------------------------------------------------------------------------
def build_backbone(name: str, freeze_ratio: float = 0.0, augment: bool = True):
"""
Operational builder:
- instantiate with DEFAULT weights
- strip classifier → features
- apply ratio-based freezing over coarse blocks
- return (model, out_dim, transform)
"""
key = (name or "").lower()
if key not in BACKBONES:
raise ValueError(
f"Unsupported backbone '{name}'. Valid options: {list_names()}"
)
spec = BACKBONES[key]
m = spec.ctor(weights=spec.weights_default)
out_dim, m = spec.strip(m)
load_backbone_weights(key, m)
mean = getattr(spec.weights_default, "meta", {}).get("mean", (0.485, 0.456, 0.406))
std = getattr(spec.weights_default, "meta", {}).get("std", (0.229, 0.224, 0.225))
crop = 299 if key == "inception_v3" else 224
if augment:
transform = transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(crop),
transforms.RandomHorizontalFlip(),
transforms.RandomVerticalFlip(),
transforms.RandomRotation(15),
transforms.ColorJitter(0.1, 0.1, 0.1, 0.05),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
]
)
else:
transform = transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(crop),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
]
)
fr = max(0.0, min(1.0, float(freeze_ratio)))
blocks = spec.blocks(m)
n = len(blocks)
freeze_n = int(math.floor(n * fr))
for b in blocks[:freeze_n]:
for p in b.parameters():
p.requires_grad = False
return m, out_dim, transform
# ---------------------------------------------------------------------------
# TowerBase ABC
# ---------------------------------------------------------------------------
class TowerBase(ABC):
"""Abstract base class for a HyperTower tower.
Concrete sub-classes must implement ``embed_dims``, ``embed_batch``, and
``prepare_fold``. Everything else has a sensible default no-op.
"""
# ------------------------------------------------------------------
# Required interface
# ------------------------------------------------------------------
@property
@abstractmethod
def embed_dims(self) -> list[int]:
"""Ordered list of embedding dimensionalities contributed to the bridge.
Most towers contribute one embedding (e.g. GeometryTower → [geom_hidden]).
ImageClinicalTower contributes two (image + clinical → [img_dim, cd_dim]).
"""
@abstractmethod
def embed_batch(
self,
batch: dict,
*,
device: torch.device,
slot: int = 1,
) -> list[torch.Tensor]:
"""Return a list of embeddings for one eye slot in *batch*.
Parameters
----------
batch : dict — batch produced by a SlotDataset loader
device : torch.device
slot : 1 (OD / image_1 / matrix_1) or 2 (OS / image_2 / matrix_2)
Returns
-------
list of Tensor — same length and order as ``embed_dims``
"""
@abstractmethod
def prepare_fold(
self,
*,
eye_train: list,
bilat_train: list,
bilat_val: list,
bilat_test: list,
image_preprocessor,
image_cache,
device: torch.device,
args,
) -> None:
"""Called once per fold before the main epoch loop."""
def early_pass(self, context: EarlyPassContext) -> None:
"""Optional: called once per fold before loaders are built."""
pass
def get_sample(self, entry) -> "torch.Tensor | dict": # noqa: ARG002
"""Optional: called by HTDataset to retrieve one sample for this tower.
entry : ShellEntry — carries entity_id, label, side, paired.
Single mode (entry.paired=False): return one Tensor.
Paired mode (entry.paired=True): return {"a": Tensor, "b": Tensor}.
Default raises NotImplementedError. Towers that participate in the
v4 HTDataset pipeline must implement this.
"""
raise NotImplementedError(
f"{type(self).__name__}.get_sample() is not implemented. "
"Implement it to use this tower with HTDataset."
)
# ------------------------------------------------------------------
# Optional interface
# ------------------------------------------------------------------
def augment_samples(self, samples: list) -> list:
"""Optional: add modality-specific keys to sample dicts before loaders are built.
Called by the orchestrator on each sample list (eye_train, bilat_train,
bilat_val, bilat_test) *after* ``prepare_fold`` completes.
The default implementation is a no-op. GeometryTower overrides this to
inject ``seg_map_1`` / ``seg_map_2`` numpy arrays so the loader can deliver
them as tensors alongside the image and clinical slots.
"""
return samples
def cd_warmup_embedding(
self,
batch: dict,
*,
device: torch.device,
) -> Optional[torch.Tensor]:
"""Return the clinical embedding for cd_warmup phase, or None if not applicable.
ClinicalDataTower overrides this to return its encoder output.
All other towers return None (the default).
This replaces isinstance(t, ClinicalDataTower) checks in train_towers_epoch,
so new tower types never require changes to towerbase.py.
"""
return None
def set_phase(self, phase: str) -> None:
"""Control requires_grad on this tower's parameters for *phase*.
Phases: ``cd_warmup``, ``tower_warmup``, ``fused_warmup``, ``main``.
Default: no-op (tower parameters always trainable unless overridden).
"""
@property
def total_epochs(self) -> int:
"""How many epochs this tower participates in the main loop."""
return 0
@property
def train_loader(self) -> Optional[DataLoader]:
"""Single-eye training loader, or None if not applicable."""
return None
@property
def cd_only_loader(self) -> Optional[DataLoader]:
"""Clinical-data-only loader for cd_warmup phase, or None."""
return None
def finalize_fold(
self,
*,
bridge,
bilat_train_loader: Optional[DataLoader] = None,
val_loader: Optional[DataLoader] = None,
device: torch.device,
args,
) -> None:
"""Optional post-epoch-loop operations (e.g. fused head training)."""
# ---------------------------------------------------------------------------
# Shared training utilities
# ---------------------------------------------------------------------------
def _set_requires_grad(module: nn.Module, enabled: bool) -> None:
for p in module.parameters():
p.requires_grad = enabled
def _to_label_tensor(labels, device: torch.device) -> torch.Tensor:
if torch.is_tensor(labels):
return labels.to(device=device, dtype=torch.long)
return torch.as_tensor(labels, dtype=torch.long, device=device)
# ---------------------------------------------------------------------------
# Modular tower training / evaluation (multi-tower interface)
# ---------------------------------------------------------------------------
def train_towers_epoch(
towers: "list[TowerBase]",
bridge: Bridge,
loader: DataLoader,
optimizer,
device: torch.device,
*,
phase: str,
bcd_prob: float = 0.5,
tower_loss_mode: str = "bcd",
) -> tuple[float, float]:
"""Train one epoch using the modular tower interface.
All towers and the bridge are set to the given phase via their
``set_phase`` methods. BCD / fused loss semantics mirror the
existing ``train_single_epoch`` logic:
cd_warmup — clinical encoder aux head only (slot index 1)
tower_warmup — img + cd aux heads equally
fused_warmup — fused bridge output only
main — BCD (randomly img or cd aux) vs fused, per tower_loss_mode
Duck typing: towers that implement ``cd_warmup_embedding`` participate in
cd_warmup; all others are skipped for that phase. No isinstance checks.
"""
for t in towers:
t.set_phase(phase)
bridge.set_phase(phase)
total_loss = total_correct = total_n = 0
for batch in loader:
y = batch.get("label_1")
if y is None:
continue
# ---- cd_warmup: train clinical encoder via its aux head ----
if phase == "cd_warmup":
z_cd = None
for t in towers:
z = t.cd_warmup_embedding(batch, device=device)
if z is not None:
z_cd = z
break
if z_cd is None:
continue
logits = bridge.aux_heads[1](z_cd)
y_t = _to_label_tensor(y, device)
loss = F.cross_entropy(logits, y_t)
optimizer.zero_grad()
loss.backward()
optimizer.step()
bs = y_t.shape[0]
total_loss += float(loss.item()) * bs
total_correct += int((logits.argmax(1) == y_t).sum())
total_n += bs
continue
# ---- collect all slot-1 embeddings from every tower ----
all_embs = []
for t in towers:
all_embs.extend(t.embed_batch(batch, device=device, slot=1))
if any(e is None for e in all_embs):
continue
y_t = _to_label_tensor(y, device)
if phase == "tower_warmup":
# Average loss across all available aux heads
aux_logits = [bridge.aux_heads[i](emb) for i, emb in enumerate(all_embs)]
loss = sum(F.cross_entropy(l, y_t) for l in aux_logits) / len(aux_logits)
# Softmax average for metrics
logits = sum(F.softmax(l, dim=1) for l in aux_logits) / len(aux_logits)
elif phase == "fused_warmup":
logits_fused, _ = bridge.fuse(all_embs)
loss = F.cross_entropy(logits_fused, y_t)
logits = logits_fused
else: # main
if tower_loss_mode == "all":
logits_fused, aux = bridge.fuse(all_embs)
loss = F.cross_entropy(logits_fused, y_t)
for aux_l in aux:
loss = loss + F.cross_entropy(aux_l, y_t)
logits = logits_fused
elif random() < bcd_prob:
# Randomly pick ONE tower to train (Generalized BCD)
idx = int(random() * len(all_embs))
logits = bridge.aux_heads[idx](all_embs[idx])
loss = F.cross_entropy(logits, y_t)
else:
logits_fused, _ = bridge.fuse(all_embs)
loss = F.cross_entropy(logits_fused, y_t)
logits = logits_fused
optimizer.zero_grad()
loss.backward()
optimizer.step()
bs = y_t.shape[0]
total_loss += float(loss.item()) * bs
total_correct += int((logits.argmax(1) == y_t).sum())
total_n += bs
return (
total_loss / total_n if total_n else float("nan"),
total_correct / total_n if total_n else float("nan"),
)
def collect_probs_towers(
towers: "list[TowerBase]",
bridge: Bridge,
loader: DataLoader,
device: torch.device,
*,
tower_mode: str = "ensemble",
) -> tuple[np.ndarray, np.ndarray]:
"""Evaluate using the modular tower interface.
Returns ``(y_true, probs)`` with patient-level probabilities
(OD + OS averaged for ensemble mode, OD-only for single mode).
"""
for t in towers:
if isinstance(t, nn.Module):
t.eval()
bridge.eval()
y_chunks, p_chunks = [], []
with torch.no_grad():
for batch in loader:
y = batch.get("label_1")
if y is None:
continue
if not (
torch.is_tensor(batch.get("image_1"))
and torch.is_tensor(batch.get("image_2"))
):
continue
all_embs_od = []
all_embs_os = []
for t in towers:
all_embs_od.extend(t.embed_batch(batch, device=device, slot=1))
all_embs_os.extend(t.embed_batch(batch, device=device, slot=2))
if any(e is None for e in all_embs_od + all_embs_os):
continue
logits_od, _ = bridge.fuse(all_embs_od)
logits_os, _ = bridge.fuse(all_embs_os)
if tower_mode == "ensemble":
probs = 0.5 * (
F.softmax(logits_od, dim=1) + F.softmax(logits_os, dim=1)
)
else:
probs = F.softmax(logits_od, dim=1)
y_chunks.append(_to_label_tensor(y, device).cpu().numpy())
p_chunks.append(probs.cpu().numpy())
for t in towers:
if isinstance(t, nn.Module):
t.train()
bridge.train()
if not y_chunks:
return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32)
return np.concatenate(y_chunks), np.concatenate(p_chunks, axis=0)
-279
View File
@@ -1,279 +0,0 @@
from __future__ import annotations
import math
from typing import Optional
import torch
from torch import nn
from torchvision import transforms
from v3.classes.backbones import BACKBONES, list_names, load_backbone_weights
from v3.classes.SE_attention import SEBlock
from v3.classes.data_bundle import DataBundle
def build_backbone(name: str, freeze_ratio: float = 0.0, augment: bool = True):
"""
Operational builder:
- instantiate with DEFAULT weights
- strip classifier → features
- apply ratio-based freezing over coarse blocks
- return (model, out_dim, transform)
"""
key = (name or "").lower()
if key not in BACKBONES:
raise ValueError(f"Unsupported backbone '{name}'. Valid options: {list_names()}")
spec = BACKBONES[key]
m = spec.ctor(weights=spec.weights_default)
out_dim, m = spec.strip(m)
load_backbone_weights(key, m)
# transforms: use the weights mean/std, but keep your augmentation pipeline
mean = getattr(spec.weights_default, "meta", {}).get("mean", (0.485, 0.456, 0.406))
std = getattr(spec.weights_default, "meta", {}).get("std", (0.229, 0.224, 0.225))
crop = 299 if key == "inception_v3" else 224
if augment:
transform = transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(crop),
transforms.RandomHorizontalFlip(),
transforms.RandomVerticalFlip(),
transforms.RandomRotation(15),
transforms.ColorJitter(0.1, 0.1, 0.1, 0.05),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
]
)
else:
transform = transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(crop),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
]
)
# ratio-based freezing: freeze earliest floor(N * freeze_ratio) blocks
fr = max(0.0, min(1.0, float(freeze_ratio)))
blocks = spec.blocks(m)
n = len(blocks)
freeze_n = int(math.floor(n * fr))
for b in blocks[:freeze_n]:
for p in b.parameters():
p.requires_grad = False
return m, out_dim, transform
class ImageTower(nn.Module):
"""
Vision backbone → pooled features.
- backbone: one of list_names() (default 'efficientnet_b0')
- always DEFAULT torchvision weights
- freeze_ratio ∈ [0,1] freezes earliest floor(N*freeze_ratio) blocks
- returns [N, out_dim] features from backbone forward
"""
def __init__(
self,
backbone: str = "efficientnet_b0",
freeze_ratio: float = 0.0,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
augment: bool = True,
geometry_dim: int = 0,
):
super().__init__()
self.backbone, base_dim, self.transform = build_backbone(
backbone, freeze_ratio, augment=augment
)
self._name = backbone
# Keep ordered blocks for dynamic freezing/thawing
key = (self._name or "").lower()
self._spec = BACKBONES[key]
self._blocks = self._spec.blocks(self.backbone)
# Optional tower-level SE over the final feature vector
self.base_dim = base_dim
self.geometry_dim = max(0, int(geometry_dim))
self.out_dim = self.base_dim + self.geometry_dim
self.tower_ln = nn.LayerNorm(self.base_dim) if se_pre_norm else nn.Identity()
self.tower_se = (
SEBlock(self.base_dim, reduction=se_reduction, residual=True)
if use_se
else None
)
def forward(
self, x: torch.Tensor, geometry: Optional[torch.Tensor] = None
) -> torch.Tensor:
y = self.backbone(x)
# sanity: pooled features, not logits
assert y.dim() == 2 and y.size(1) == self.base_dim, (
f"Expected features [N,{self.base_dim}], got {tuple(y.shape)}"
)
if self.tower_se is not None:
y, _ = self.tower_se(self.tower_ln(y))
if self.geometry_dim > 0:
if geometry is None or geometry.numel() == 0:
geom = torch.zeros(
y.size(0), self.geometry_dim, device=y.device, dtype=y.dtype
)
else:
if geometry.dim() == 1:
geom = geometry.unsqueeze(0)
else:
geom = geometry
geom = geom.to(device=y.device, dtype=y.dtype)
if geom.size(0) != y.size(0):
raise ValueError(
f"Geometry batch size mismatch: {geom.size(0)} vs {y.size(0)}"
)
if geom.size(1) != self.geometry_dim:
raise ValueError(
f"Expected geometry dim {self.geometry_dim}, got {geom.size(1)}"
)
y = torch.cat([y, geom], dim=1)
return y
def set_freeze_ratio(self, ratio: float):
"""Dynamically freeze earliest floor(N*ratio) backbone blocks."""
r = max(0.0, min(1.0, float(ratio)))
n = len(self._blocks)
freeze_n = int(math.floor(n * r))
# Unfreeze all first
for b in self._blocks:
for p in b.parameters():
p.requires_grad = True
# Freeze earliest blocks
for b in self._blocks[:freeze_n]:
for p in b.parameters():
p.requires_grad = False
class SiameseImageTower(nn.Module):
"""
Shared-weight bilateral image tower.
Runs OD and OS images through a single shared backbone, then returns
cat([f_mean, f_delta]) where:
f_mean = (f_od + f_os) / 2 -- shared bilateral representation
f_delta = f_od - f_os -- asymmetry, signed OD-relative
out_dim = 2 * backbone_out_dim
When x_os is None (single-eye fallback):
f_mean = f_od
f_delta = zeros
so the module degrades gracefully when only one eye is available.
The shared backbone means both eyes contribute to every gradient update,
effectively doubling the training signal for the visual pathway without
doubling parameters.
"""
def __init__(
self,
backbone: str = "efficientnet_b0",
freeze_ratio: float = 0.0,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
augment: bool = True,
):
super().__init__()
self._tower = ImageTower(
backbone=backbone,
freeze_ratio=freeze_ratio,
use_se=use_se,
se_reduction=se_reduction,
se_pre_norm=se_pre_norm,
augment=augment,
geometry_dim=0,
)
self.out_dim = self._tower.out_dim * 2
self.transform = self._tower.transform
def forward(
self,
x_od: torch.Tensor,
x_os: Optional[torch.Tensor] = None,
) -> torch.Tensor:
f_od = self._tower(x_od)
if x_os is None:
f_mean = f_od
f_delta = torch.zeros_like(f_od)
else:
f_os = self._tower(x_os)
f_mean = (f_od + f_os) * 0.5
f_delta = f_od - f_os
return torch.cat([f_mean, f_delta], dim=1)
def set_freeze_ratio(self, ratio: float) -> None:
"""Delegates to the shared inner tower."""
self._tower.set_freeze_ratio(ratio)
class ClinicalTower(nn.Module):
"""MLP over DataBundle.vectorize_row outputs (convert to torch inside tower)."""
def __init__(
self,
clinical_data: DataBundle,
hidden_dim: int = 128,
dropout: float = 0.1,
use_se: bool = False,
se_reduction: int = 16,
se_pre_norm: bool = True,
):
super().__init__()
self.feature_dim = clinical_data.feature_dim
self.out_dim = hidden_dim
# two-block MLP so we can optionally freeze/thaw per block
self.block0 = nn.Sequential(
nn.Linear(self.feature_dim, hidden_dim),
nn.LayerNorm(hidden_dim),
nn.ReLU(inplace=True),
nn.Dropout(dropout),
)
self.block1 = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(inplace=True),
)
self.net = nn.Sequential(self.block0, self.block1)
self.tower_ln = nn.LayerNorm(hidden_dim) if se_pre_norm else nn.Identity()
self.tower_se = (
SEBlock(hidden_dim, reduction=se_reduction, residual=True)
if use_se
else None
)
def forward(self, meta_np_or_torch) -> torch.Tensor:
if isinstance(meta_np_or_torch, torch.Tensor):
x = meta_np_or_torch
else:
x = torch.as_tensor(meta_np_or_torch, dtype=torch.float32)
h = self.net(x)
if self.tower_se is not None:
h, _ = self.tower_se(self.tower_ln(h))
return h
def set_freeze_ratio(self, ratio: float):
"""Optionally freeze earliest blocks of the MLP."""
r = max(0.0, min(1.0, float(ratio)))
# Unfreeze all
for p in self.block0.parameters():
p.requires_grad = True
for p in self.block1.parameters():
p.requires_grad = True
# Freeze earliest blocks based on ratio threshold
if r >= 0.5:
for p in self.block0.parameters():
p.requires_grad = False
if r >= 1.0:
for p in self.block1.parameters():
p.requires_grad = False
File diff suppressed because it is too large Load Diff
+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")