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