"""hypertower_models — HyperTower vehicle classes and their training/eval helpers. A "vehicle" wires one or more tower encoders together with a Bridge to form a complete trainable model. Vehicles can in principle be run standalone; the v3_hypertower orchestrator drives them through the full fold/epoch loop. Contents -------- SingleEyeHT — ImageEncoder + ClinicalEncoder + Bridge (eye-level) BilateralHT — shared eye towers + joint fusion layers + Bridge SiameseHT — SiameseImageTower + Bridge (image-only bilateral) FusedEnsembleHT — SingleEyeHT base + per-eye attention scorer LogitMLPEnsembleHT — MLP head over concatenated per-eye logits EmbeddingMLPEnsembleHT — MLP head over concatenated per-eye z_fused embeddings NTowerHT — N named encoders + Bridge (general, key-mapped) NLateralHT — N same-type inputs through shared encoder + joint MLP MonoTowerHT — single tower + direct classifier (no bridge) Training helpers : train_single_epoch, train_bilateral_epoch, train_siamese_epoch, train_fusion_epoch, train_ntower_epoch, train_mono_epoch Inference helpers : collect_probs_classic, collect_probs_ensemble, collect_probs_ensemble_pereye, collect_probs_bilateral, collect_probs_bilateral_components, collect_probs_siamese, collect_probs_fused, collect_probs_single_components, collect_probs_eye_level, collect_probs_ntower, collect_probs_mono """ from __future__ import annotations from random import random 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 v3.classes.towerbase import _set_requires_grad, _to_label_tensor from v3.classes.image_towers import ImageEncoder, SiameseImageTower from v3.classes.clinical_towers import ClinicalEncoder from v3.classes.bridges import Bridge, HTClassifier, HyperBridge # --------------------------------------------------------------------------- # Vehicle classes # --------------------------------------------------------------------------- class SingleEyeHT(nn.Module): """ImageEncoder + ClinicalEncoder + Bridge, trained on eye-level samples.""" def __init__( self, *, backbone: str, freeze_ratio: float, augment: bool, clinical_data, num_classes: int, cd_hidden_dim: int = 128, fusion_dim: int = 256, bridge_mode: str = "fused", bridge_dropout: float = 0.5, cd_dropout: float = 0.1, se_img_tower: bool = False, se_cd_tower: bool = False, se_bridge: bool = False, ): super().__init__() self.img_tower = ImageEncoder( backbone=backbone, freeze_ratio=freeze_ratio, augment=augment, use_se=se_img_tower, ) self.cd_tower = ClinicalEncoder( clinical_data=clinical_data, hidden_dim=cd_hidden_dim, dropout=cd_dropout, use_se=se_cd_tower, ) self.bridge = Bridge( tower_dims=[self.img_tower.out_dim, self.cd_tower.out_dim], num_classes=num_classes, fusion_dim=fusion_dim, mode=bridge_mode, dropout=bridge_dropout, use_se=se_bridge, ) @property def transform(self): return self.img_tower.transform def encode(self, x: torch.Tensor, meta: torch.Tensor) -> torch.Tensor: """Return z_fused (fusion_dim) without the classifier head.""" return self.bridge.encode([self.img_tower(x), self.cd_tower(meta)]) def forward(self, x: torch.Tensor, meta: torch.Tensor) -> torch.Tensor: img_feats = None if self.bridge.mode == "clinical_only" else self.img_tower(x) md_feats = None if self.bridge.mode == "image_only" else self.cd_tower(meta) out_f, _ = self.bridge.fuse([img_feats, md_feats]) return out_f class NTowerHT(nn.Module): """General N-tower vehicle: any named encoder modules fused through a Bridge. Parameters ---------- towers : ordered dict of {name: encoder_module}. Each module must expose ``.out_dim``. Bridge slot order follows dict order. num_classes : number of output classes fusion_dim : projection dimensionality inside the bridge dropout : dropout before the fused output head use_se : SE gate on the fused vector Forward contract ---------------- ``forward(embeddings)`` takes a ``dict[str, Tensor]`` of pre-computed per-tower embeddings (keyed by tower name) and returns ``(logits_fused, aux_dict)`` where ``aux_dict`` maps each tower name to its auxiliary head logits. The vehicle does not know how to extract embeddings from raw data — that is the driver's responsibility. Use ``embed_batch`` (TowerBase API) or custom extraction logic in the training loop, then pass the result here. Type-aware helpers ------------------ ``transform`` — returns the ``.transform`` of the first tower that has one (typically the image tower), for use by data loaders. """ def __init__( self, towers: dict[str, nn.Module], num_classes: int, fusion_dim: int = 256, dropout: float = 0.5, use_se: bool = False, ): super().__init__() self.towers = nn.ModuleDict(towers) self.bridge = Bridge( tower_dims=[t.out_dim for t in self.towers.values()], num_classes=num_classes, fusion_dim=fusion_dim, dropout=dropout, use_se=use_se, ) @property def transform(self): """Image transform from the first tower that exposes one, or None.""" for t in self.towers.values(): if hasattr(t, "transform"): return t.transform return None def encode(self, embeddings: dict[str, torch.Tensor]) -> torch.Tensor: """Return z_fused (pre-classifier) from a dict of per-tower embeddings.""" return self.bridge.encode([embeddings[name] for name in self.towers]) def forward( self, embeddings: dict[str, torch.Tensor], ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Fuse pre-computed embeddings through the bridge. Returns ------- logits_fused : Tensor [B, num_classes] aux_logits : dict mapping tower name → Tensor [B, num_classes] """ ordered = [embeddings[name] for name in self.towers] logits_fused, aux = self.bridge.fuse(ordered) return logits_fused, {name: aux[i] for i, name in enumerate(self.towers)} class NLateralHT(nn.Module): """N same-type inputs through a shared encoder, jointly compressed, then classified. Intended for encoding multiple instances of the same modality together — e.g. left + right fundus images, or OD + OS clinical vectors. All inputs share the same encoder weights (one forward pass per input). Output matches NTowerHT's contract: ``(logits, {name: aux_logits})``. Aux logits are per-input classifications taken before the joint MLP, useful for BCD-style training. Parameters ---------- encoder : shared encoder module with ``.out_dim`` input_names : ordered names for each input slot (e.g. ``["od", "os"]``) num_classes : output classes fusion_dim : hidden dim of the joint MLP dropout : dropout in joint MLP and classifier head """ def __init__( self, encoder: nn.Module, input_names: list[str], num_classes: int, fusion_dim: int = 256, dropout: float = 0.5, ): super().__init__() self.encoder = encoder self.input_names = list(input_names) n = len(input_names) in_dim: int = encoder.out_dim # type: ignore[assignment] # Joint MLP: [z0 ‖ z1 ‖ ... ‖ z_{n-1}] → fusion_dim → in_dim self.joint = nn.Sequential( nn.Linear(n * in_dim, fusion_dim), nn.LayerNorm(fusion_dim), nn.ReLU(), nn.Dropout(dropout), nn.Linear(fusion_dim, in_dim), ) # Per-input aux heads — classify each input before joining self.aux_heads = nn.ModuleList([ nn.Linear(in_dim, num_classes) for _ in range(n) ]) # Main classifier on the joint representation self.head = nn.Sequential( nn.ReLU(), nn.Dropout(dropout), nn.Linear(in_dim, num_classes), ) @property def transform(self): return getattr(self.encoder, "transform", None) def encode(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: """Return joint embedding (post-MLP, pre-classifier).""" zs = [self.encoder(inputs[name]) for name in self.input_names] return self.joint(torch.cat(zs, dim=1)) def forward( self, inputs: dict[str, torch.Tensor], ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Encode inputs jointly and classify. Parameters ---------- inputs : ``{name: raw_tensor}`` for each input slot Returns ------- logits : [B, num_classes] aux_dict : ``{name: [B, num_classes]}`` — per-input pre-join logits """ zs = [self.encoder(inputs[name]) for name in self.input_names] z_joint = self.joint(torch.cat(zs, dim=1)) logits = self.head(z_joint) aux = {name: head(z) for name, head, z in zip(self.input_names, self.aux_heads, zs)} return logits, aux class BilateralHT(nn.Module): """Bilateral vehicle: shared eye-level towers + joint fusion layers + Bridge.""" def __init__( self, *, backbone: str, freeze_ratio: float, augment: bool, clinical_data, num_classes: int, cd_hidden_dim: int = 128, fusion_dim: int = 256, ): super().__init__() self.eye_img_tower = ImageEncoder( backbone=backbone, freeze_ratio=freeze_ratio, augment=augment, use_se=False, ) self.eye_cd_tower = ClinicalEncoder( clinical_data=clinical_data, hidden_dim=cd_hidden_dim, use_se=False, ) img_dim = self.eye_img_tower.out_dim md_dim = self.eye_cd_tower.out_dim self.joint_img = nn.Sequential( nn.Linear(2 * img_dim, fusion_dim), nn.LayerNorm(fusion_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(fusion_dim, img_dim), ) self.joint_md = nn.Sequential( nn.Linear(2 * md_dim, fusion_dim), nn.LayerNorm(fusion_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(fusion_dim, md_dim), ) self.bridge = Bridge(tower_dims=[img_dim, md_dim], num_classes=num_classes, fusion_dim=fusion_dim, mode="fused", use_se=False) self.aux_img = nn.Linear(img_dim, num_classes) self.aux_md = nn.Linear(md_dim, num_classes) @property def transform(self): return self.eye_img_tower.transform def encode_joint( self, x_od: torch.Tensor, meta_od: torch.Tensor, x_os: torch.Tensor, meta_os: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: img_od = self.eye_img_tower(x_od); md_od = self.eye_cd_tower(meta_od) img_os = self.eye_img_tower(x_os); md_os = self.eye_cd_tower(meta_os) return (self.joint_img(torch.cat([img_od, img_os], dim=1)), self.joint_md(torch.cat([md_od, md_os], dim=1))) def forward( self, x_od: torch.Tensor, meta_od: torch.Tensor, x_os: torch.Tensor, meta_os: torch.Tensor, ) -> torch.Tensor: joint_img, joint_md = self.encode_joint(x_od, meta_od, x_os, meta_os) out_f, _ = self.bridge.fuse([joint_img, joint_md]) return out_f class SiameseHT(nn.Module): """Bilateral vehicle using a shared-weight SiameseImageTower (mean+delta).""" def __init__( self, *, backbone: str, freeze_ratio: float, augment: bool, num_classes: int, fusion_dim: int = 256, ): super().__init__() self.img_tower = SiameseImageTower( backbone=backbone, freeze_ratio=freeze_ratio, augment=augment, use_se=False, ) img_dim = self.img_tower.out_dim self.bridge = Bridge(tower_dims=[img_dim], num_classes=num_classes, fusion_dim=fusion_dim, use_se=False) self.aux_img = nn.Linear(img_dim, num_classes) @property def transform(self): return self.img_tower.transform def encode(self, x_od: torch.Tensor, x_os: torch.Tensor) -> torch.Tensor: return self.img_tower(x_od, x_os) def forward(self, x_od: torch.Tensor, x_os: torch.Tensor) -> torch.Tensor: out_f, _ = self.bridge.fuse([self.encode(x_od, x_os)]) return out_f class FusedEnsembleHT(nn.Module): """SingleEyeHT base with a per-eye attention scorer for bilateral fusion.""" def __init__(self, base: SingleEyeHT, num_classes: int): super().__init__() self.base = base self.eye_scorer = nn.Linear(num_classes, 1, bias=True) @property def head(self) -> nn.Module: return self.eye_scorer def forward( self, x_od: torch.Tensor, meta_od: torch.Tensor, x_os: torch.Tensor, meta_os: torch.Tensor, ) -> torch.Tensor: logit_od = self.base(x_od, meta_od) logit_os = self.base(x_os, meta_os) scores = torch.cat([self.eye_scorer(logit_od), self.eye_scorer(logit_os)], dim=1) alpha = torch.softmax(scores, dim=1) return alpha[:, 0:1] * logit_od + alpha[:, 1:2] * logit_os class LogitMLPEnsembleHT(nn.Module): """MLP head trained on concatenated per-eye logits.""" def __init__(self, base: SingleEyeHT, num_classes: int, hidden: int = 64): super().__init__() self.base = base self.head = nn.Sequential( nn.Linear(2 * num_classes, hidden), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden, num_classes), ) def forward( self, x_od: torch.Tensor, meta_od: torch.Tensor, x_os: torch.Tensor, meta_os: torch.Tensor, ) -> torch.Tensor: return self.head(torch.cat([self.base(x_od, meta_od), self.base(x_os, meta_os)], dim=1)) class EmbeddingMLPEnsembleHT(nn.Module): """MLP head trained on concatenated per-eye z_fused embeddings.""" def __init__(self, base: SingleEyeHT, num_classes: int, hidden: int = 256): super().__init__() self.base = base fusion_dim = base.bridge.W[0].out_features self.head = nn.Sequential( nn.Linear(2 * fusion_dim, hidden), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden, num_classes), ) def forward( self, x_od: torch.Tensor, meta_od: torch.Tensor, x_os: torch.Tensor, meta_os: torch.Tensor, ) -> torch.Tensor: return self.head(torch.cat([self.base.encode(x_od, meta_od), self.base.encode(x_os, meta_os)], dim=1)) # --------------------------------------------------------------------------- # MonoTowerHT — single tower + Bridge(N=1) # --------------------------------------------------------------------------- class MonoTowerHT(nn.Module): """Single-tower vehicle: tower output fed directly into a classifier head. No bridge projection — the tower's embedding goes straight to ``ReLU → Dropout → Linear(out_dim → num_classes)``. This is the minimal architecture: just the tower's learned representation with a classification head attached. Contrast with ``NTowerHT(N=1)``, which still projects through the bridge's shared ``fusion_dim`` space. These are distinct architectures and may yield different results. Parameters ---------- tower : any nn.Module with ``.out_dim`` (ImageEncoder, ClinicalEncoder, …) num_classes : number of output classes dropout : dropout before the output linear layer """ def __init__( self, tower: nn.Module, num_classes: int, dropout: float = 0.5, ): super().__init__() self.tower = tower out_dim: int = tower.out_dim # type: ignore[assignment] self.head = nn.Sequential( nn.ReLU(), nn.Dropout(dropout), nn.Linear(out_dim, num_classes), ) def classify(self, z: torch.Tensor) -> torch.Tensor: """Classify a pre-computed embedding.""" return self.head(z) def forward(self, *args, **kwargs) -> torch.Tensor: """Pass inputs through tower then classifier head.""" return self.head(self.tower(*args, **kwargs)) # --------------------------------------------------------------------------- # Phase control # --------------------------------------------------------------------------- def _set_single_phase(model: SingleEyeHT, phase: str) -> None: bridge_mode = model.bridge.mode if bridge_mode in ("image_only", "clinical_only") and phase == "fused_warmup": phase = "tower_warmup" if phase == "cd_warmup": _set_requires_grad(model.img_tower, False) _set_requires_grad(model.cd_tower, True) _set_requires_grad(model.bridge.aux_heads[0], False) _set_requires_grad(model.bridge.aux_heads[1], True) _set_requires_grad(model.bridge.W[0], False) _set_requires_grad(model.bridge.W[1], False) _set_requires_grad(model.bridge.classifier_fused, False) return if phase == "tower_warmup": _set_requires_grad(model.img_tower, bridge_mode != "clinical_only") _set_requires_grad(model.cd_tower, bridge_mode != "image_only") _set_requires_grad(model.bridge.aux_heads[0], bridge_mode != "clinical_only") _set_requires_grad(model.bridge.aux_heads[1], bridge_mode != "image_only") _set_requires_grad(model.bridge.W[0], False) _set_requires_grad(model.bridge.W[1], False) _set_requires_grad(model.bridge.classifier_fused, False) return if phase == "fused_warmup": _set_requires_grad(model.img_tower, False) _set_requires_grad(model.cd_tower, False) _set_requires_grad(model.bridge.aux_heads[0], False) _set_requires_grad(model.bridge.aux_heads[1], False) _set_requires_grad(model.bridge.W[0], True) _set_requires_grad(model.bridge.W[1], True) _set_requires_grad(model.bridge.classifier_fused, True) return _set_requires_grad(model, True) def _set_bilateral_phase(model: BilateralHT, phase: str) -> None: if phase == "tower_warmup": for m in (model.eye_img_tower, model.eye_cd_tower, model.joint_img, model.joint_md, model.aux_img, model.aux_md): _set_requires_grad(m, True) _set_requires_grad(model.bridge, False) return if phase == "fused_warmup": for m in (model.eye_img_tower, model.eye_cd_tower, model.joint_img, model.joint_md, model.aux_img, model.aux_md): _set_requires_grad(m, False) _set_requires_grad(model.bridge, True) return _set_requires_grad(model, True) # --------------------------------------------------------------------------- # Training helpers # --------------------------------------------------------------------------- def train_single_epoch( model: SingleEyeHT, loader: DataLoader, opt, device: torch.device, *, phase: str, bcd_prob: float = 0.5, tower_loss_mode: str = "bcd", ) -> tuple[float, float]: model.train() _set_single_phase(model, phase) total_loss = total_correct = total_n = 0 for batch in loader: x = batch.get("image_1"); m = batch.get("matrix_1"); y = batch.get("label_1") if phase == "cd_warmup": if not torch.is_tensor(m): continue y = _to_label_tensor(y, device) logits = model.bridge.aux_heads[1](model.cd_tower(m.to(device))) loss = F.cross_entropy(logits, y) opt.zero_grad(); loss.backward(); opt.step() bs = y.shape[0] total_loss += float(loss.item()) * bs total_correct += int((logits.argmax(1) == y).sum()) total_n += bs continue if not (torch.is_tensor(x) and torch.is_tensor(m)): continue x = x.to(device); m = m.to(device); y = _to_label_tensor(y, device) bridge_mode = model.bridge.mode img_feats = None if bridge_mode == "clinical_only" else model.img_tower(x) md_feats = None if bridge_mode == "image_only" else model.cd_tower(m) if phase == "tower_warmup": if bridge_mode == "clinical_only": logits = model.bridge.aux_heads[1](md_feats); loss = F.cross_entropy(logits, y) elif bridge_mode == "image_only": logits = model.bridge.aux_heads[0](img_feats); loss = F.cross_entropy(logits, y) else: li = model.bridge.aux_heads[0](img_feats); lm = model.bridge.aux_heads[1](md_feats) loss = 0.5 * (F.cross_entropy(li, y) + F.cross_entropy(lm, y)) logits = 0.5 * (F.softmax(li, dim=1) + F.softmax(lm, dim=1)) elif phase == "fused_warmup": logits, _ = model.bridge.fuse([img_feats, md_feats]); loss = F.cross_entropy(logits, y) else: if bridge_mode == "clinical_only": logits = model.bridge.aux_heads[1](md_feats); loss = F.cross_entropy(logits, y) elif bridge_mode == "image_only": logits = model.bridge.aux_heads[0](img_feats); loss = F.cross_entropy(logits, y) elif tower_loss_mode == "all": li = model.bridge.aux_heads[0](img_feats); lm = model.bridge.aux_heads[1](md_feats) logits, _ = model.bridge.fuse([img_feats, md_feats]) loss = F.cross_entropy(logits, y) + F.cross_entropy(li, y) + F.cross_entropy(lm, y) elif random() < bcd_prob: logits = (model.bridge.aux_heads[0](img_feats) if random() < 0.5 else model.bridge.aux_heads[1](md_feats)) loss = F.cross_entropy(logits, y) else: logits, _ = model.bridge.fuse([img_feats, md_feats]); loss = F.cross_entropy(logits, y) opt.zero_grad(); loss.backward(); opt.step() bs = y.shape[0] total_loss += float(loss.item()) * bs total_correct += int((logits.argmax(1) == y).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 train_bilateral_epoch( model: BilateralHT, loader: DataLoader, opt, device: torch.device, *, phase: str, bcd_prob: float = 0.5, tower_loss_mode: str = "bcd", ) -> tuple[float, float]: model.train() _set_bilateral_phase(model, phase) total_loss = total_correct = total_n = 0 for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue x1 = x1.to(device); m1 = m1.to(device) x2 = x2.to(device); m2 = m2.to(device); y = _to_label_tensor(y, device) joint_img, joint_md = model.encode_joint(x1, m1, x2, m2) if phase == "tower_warmup": li = model.aux_img(joint_img); lm = model.aux_md(joint_md) loss = 0.5 * (F.cross_entropy(li, y) + F.cross_entropy(lm, y)) logits = 0.5 * (F.softmax(li, dim=1) + F.softmax(lm, dim=1)) elif phase == "fused_warmup": logits, _ = model.bridge.fuse([joint_img, joint_md]); loss = F.cross_entropy(logits, y) else: if tower_loss_mode == "all": li = model.aux_img(joint_img); lm = model.aux_md(joint_md) logits, _ = model.bridge.fuse([joint_img, joint_md]) loss = F.cross_entropy(logits, y) + F.cross_entropy(li, y) + F.cross_entropy(lm, y) elif random() < bcd_prob: logits = model.aux_img(joint_img) if random() < 0.5 else model.aux_md(joint_md) loss = F.cross_entropy(logits, y) else: logits, _ = model.bridge.fuse([joint_img, joint_md]); loss = F.cross_entropy(logits, y) opt.zero_grad(); loss.backward(); opt.step() bs = y.shape[0] total_loss += float(loss.item()) * bs total_correct += int((logits.argmax(1) == y).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 train_siamese_epoch( model: SiameseHT, loader: DataLoader, opt, device: torch.device, *, bcd_prob: float = 0.5, tower_loss_mode: str = "bcd", ) -> tuple[float, float]: model.train() total_loss = total_correct = total_n = 0 for batch in loader: x1 = batch.get("image_1"); x2 = batch.get("image_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(x2)): continue x1 = x1.to(device); x2 = x2.to(device); y = _to_label_tensor(y, device) feats = model.encode(x1, x2) if tower_loss_mode == "all": out_f, _ = model.bridge.fuse([feats, None]) loss = F.cross_entropy(out_f, y) + F.cross_entropy(model.aux_img(feats), y) logits = out_f elif random() < bcd_prob: logits = model.aux_img(feats); loss = F.cross_entropy(logits, y) else: logits, _ = model.bridge.fuse([feats, None]); loss = F.cross_entropy(logits, y) opt.zero_grad(); loss.backward(); opt.step() bs = y.shape[0] total_loss += float(loss.item()) * bs total_correct += int((logits.argmax(1) == y).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 train_fusion_epoch( model, # FusedEnsembleHT | LogitMLPEnsembleHT | EmbeddingMLPEnsembleHT loader: DataLoader, opt, device: torch.device, ) -> tuple[float, float]: """Train only the fusion head; base SingleEyeHT is frozen in eval mode.""" model.base.eval(); model.head.train() total_loss = total_correct = total_n = 0 for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) out = model(x1.to(device), m1.to(device), x2.to(device), m2.to(device)) loss = F.cross_entropy(out, y_t) opt.zero_grad(); loss.backward(); opt.step() bs = y_t.shape[0] total_loss += float(loss.item()) * bs total_correct += int((out.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")) # --------------------------------------------------------------------------- # Inference helpers # --------------------------------------------------------------------------- def collect_probs_classic( model: SingleEyeHT, loader: DataLoader, device: torch.device, ) -> tuple[np.ndarray, np.ndarray]: model.eval() y_c, p_c = [], [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) p_od = F.softmax(model(x1.to(device), m1.to(device)), dim=1) p_os = F.softmax(model(x2.to(device), m2.to(device)), dim=1) y_np = y_t.cpu().numpy() y_c += [y_np, y_np]; p_c += [p_od.cpu().numpy(), p_os.cpu().numpy()] if not y_c: return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) return np.concatenate(y_c), np.concatenate(p_c, axis=0) def collect_probs_ensemble( model: SingleEyeHT, loader: DataLoader, device: torch.device, ) -> tuple[np.ndarray, np.ndarray]: model.eval() y_c, p_c = [], [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) p_od = F.softmax(model(x1.to(device), m1.to(device)), dim=1) p_os = F.softmax(model(x2.to(device), m2.to(device)), dim=1) y_c.append(y_t.cpu().numpy()); p_c.append((0.5 * (p_od + p_os)).cpu().numpy()) if not y_c: return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) return np.concatenate(y_c), np.concatenate(p_c, axis=0) def collect_probs_ensemble_pereye( model: SingleEyeHT, loader: DataLoader, device: torch.device, *, return_ids: bool = False, ): model.eval() y_c = []; pf_od_c, pi_od_c, pm_od_c = [], [], []; pf_os_c, pi_os_c, pm_os_c = [], [], [] id_c: list = [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) def _fwd(x, m): img_f = None if model.bridge.mode == "clinical_only" else model.img_tower(x.to(device)) md_f = None if model.bridge.mode == "image_only" else model.cd_tower(m.to(device)) out_f, aux = model.bridge.fuse([img_f, md_f]) out_i, out_m = aux[0], aux[1] pf = F.softmax(out_f, dim=1) pi = F.softmax(out_i, dim=1) if out_i is not None else pf pm = F.softmax(out_m, dim=1) if out_m is not None else pf return pf, pi, pm pf_od, pi_od, pm_od = _fwd(x1, m1); pf_os, pi_os, pm_os = _fwd(x2, m2) y_c.append(y_t.cpu().numpy()) for lst, t in ((pf_od_c, pf_od), (pi_od_c, pi_od), (pm_od_c, pm_od), (pf_os_c, pf_os), (pi_os_c, pi_os), (pm_os_c, pm_os)): lst.append(t.cpu().numpy()) if return_ids: ids = batch.get("id_1", [""] * len(y_t)) id_c.extend([str(i) for i in (ids.tolist() if torch.is_tensor(ids) else ids)]) if not y_c: z = np.zeros((0, 0), dtype=np.float32); ei = np.array([], dtype=np.int64) base = (ei, z, z, z, z, z, z) return base + (np.array([], dtype=object),) if return_ids else base y = np.concatenate(y_c) pf_od, pi_od, pm_od = (np.concatenate(c, axis=0) for c in (pf_od_c, pi_od_c, pm_od_c)) pf_os, pi_os, pm_os = (np.concatenate(c, axis=0) for c in (pf_os_c, pi_os_c, pm_os_c)) if return_ids: return y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os, np.array(id_c, dtype=object) return y, pf_od, pi_od, pm_od, pf_os, pi_os, pm_os def collect_probs_bilateral( model: BilateralHT, loader: DataLoader, device: torch.device, ) -> tuple[np.ndarray, np.ndarray]: model.eval() y_c, p_c = [], [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) p = F.softmax(model(x1.to(device), m1.to(device), x2.to(device), m2.to(device)), dim=1) y_c.append(y_t.cpu().numpy()); p_c.append(p.cpu().numpy()) if not y_c: return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) return np.concatenate(y_c), np.concatenate(p_c, axis=0) def collect_probs_bilateral_components( model: BilateralHT, loader: DataLoader, device: torch.device, ) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: model.eval() y_c = []; pf_c, pi_c, pm_c = [], [], [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) ji, jm = model.encode_joint(x1.to(device), m1.to(device), x2.to(device), m2.to(device)) out_f, _ = model.bridge.fuse([ji, jm]) y_c.append(y_t.cpu().numpy()) pf_c.append(F.softmax(out_f, dim=1).cpu().numpy()) pi_c.append(F.softmax(model.aux_img(ji), dim=1).cpu().numpy()) pm_c.append(F.softmax(model.aux_md(jm), dim=1).cpu().numpy()) if not y_c: z = np.zeros((0, 0), dtype=np.float32) return np.array([], dtype=np.int64), z, z, z return (np.concatenate(y_c), np.concatenate(pf_c, axis=0), np.concatenate(pi_c, axis=0), np.concatenate(pm_c, axis=0)) def collect_probs_siamese( model: SiameseHT, loader: DataLoader, device: torch.device, ) -> tuple[np.ndarray, np.ndarray]: model.eval() y_c, p_c = [], [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); x2 = batch.get("image_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(x2)): continue feats = model.encode(x1.to(device), x2.to(device)) logits, _ = model.bridge.fuse([feats]) y_c.append(np.array(y) if not torch.is_tensor(y) else y.cpu().numpy()) p_c.append(torch.softmax(logits, dim=1).cpu().numpy()) return np.concatenate(y_c, axis=0), np.concatenate(p_c, axis=0) def collect_probs_fused( model: FusedEnsembleHT, loader: DataLoader, device: torch.device, ) -> tuple[np.ndarray, np.ndarray]: model.eval() y_c, p_c = [], [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) p = F.softmax(model(x1.to(device), m1.to(device), x2.to(device), m2.to(device)), dim=1) y_c.append(y_t.cpu().numpy()); p_c.append(p.cpu().numpy()) if not y_c: return np.array([], dtype=np.int64), np.zeros((0, 0), dtype=np.float32) return np.concatenate(y_c), np.concatenate(p_c, axis=0) def collect_probs_single_components( model: SingleEyeHT, loader: DataLoader, device: torch.device, *, aggregate_patient: bool, return_logits: bool = False, ): model.eval() y_c = []; pf_c, pi_c, pm_c = [], [], []; lf_c, li_c, lm_c = [], [], [] with torch.no_grad(): for batch in loader: x1 = batch.get("image_1"); m1 = batch.get("matrix_1") x2 = batch.get("image_2"); m2 = batch.get("matrix_2"); y = batch.get("label_1") if not (torch.is_tensor(x1) and torch.is_tensor(m1) and torch.is_tensor(x2) and torch.is_tensor(m2)): continue y_t = _to_label_tensor(y, device) def _per_eye(x, m): img_f = None if model.bridge.mode == "clinical_only" else model.img_tower(x.to(device)) md_f = None if model.bridge.mode == "image_only" else model.cd_tower(m.to(device)) out_f, aux = model.bridge.fuse([img_f, md_f]) out_i, out_m = aux[0], aux[1] pf = F.softmax(out_f, dim=1) pi = F.softmax(out_i, dim=1) if out_i is not None else pf pm = F.softmax(out_m, dim=1) if out_m is not None else pf return pf, pi, pm, out_f, (out_i if out_i is not None else out_f), (out_m if out_m is not None else out_f) pf_od, pi_od, pm_od, lf_od, li_od, lm_od = _per_eye(x1, m1) pf_os, pi_os, pm_os, lf_os, li_os, lm_os = _per_eye(x2, m2) if aggregate_patient: y_c.append(y_t.cpu().numpy()) pf_c.append((0.5*(pf_od+pf_os)).cpu().numpy()) pi_c.append((0.5*(pi_od+pi_os)).cpu().numpy()) pm_c.append((0.5*(pm_od+pm_os)).cpu().numpy()) lf_c.append((0.5*(lf_od+lf_os)).cpu().numpy()) li_c.append((0.5*(li_od+li_os)).cpu().numpy()) lm_c.append((0.5*(lm_od+lm_os)).cpu().numpy()) else: y_np = y_t.cpu().numpy(); y_c += [y_np, y_np] pf_c += [pf_od.cpu().numpy(), pf_os.cpu().numpy()] pi_c += [pi_od.cpu().numpy(), pi_os.cpu().numpy()] pm_c += [pm_od.cpu().numpy(), pm_os.cpu().numpy()] lf_c += [lf_od.cpu().numpy(), lf_os.cpu().numpy()] li_c += [li_od.cpu().numpy(), li_os.cpu().numpy()] lm_c += [lm_od.cpu().numpy(), lm_os.cpu().numpy()] if not y_c: z = np.zeros((0, 0), dtype=np.float32) if return_logits: return np.array([], dtype=np.int64), z, z, z, z, z, z return np.array([], dtype=np.int64), z, z, z y = np.concatenate(y_c) pf, pi, pm = (np.concatenate(c, axis=0) for c in (pf_c, pi_c, pm_c)) if return_logits: lf, li, lm = (np.concatenate(c, axis=0) for c in (lf_c, li_c, lm_c)) return y, pf, pi, pm, lf, li, lm return y, pf, pi, pm def collect_probs_eye_level( model: SingleEyeHT, loader: DataLoader, device: torch.device, *, return_ids: bool = False, ): model.eval() y_c, pf_c, pi_c, pm_c, id_c = [], [], [], [], [] with torch.no_grad(): for batch in loader: x = batch.get("image_1"); m = batch.get("matrix_1"); y = batch.get("label_1") if not (torch.is_tensor(x) and torch.is_tensor(m)): continue y_t = _to_label_tensor(y, device) img_f = None if model.bridge.mode == "clinical_only" else model.img_tower(x.to(device)) md_f = None if model.bridge.mode == "image_only" else model.cd_tower(m.to(device)) out_f, aux = model.bridge.fuse([img_f, md_f]) out_i, out_m = aux[0], aux[1] pf = F.softmax(out_f, dim=1) pi = F.softmax(out_i, dim=1) if out_i is not None else pf pm = F.softmax(out_m, dim=1) if out_m is not None else pf y_c.append(y_t.cpu().numpy()) pf_c.append(pf.cpu().numpy()); pi_c.append(pi.cpu().numpy()); pm_c.append(pm.cpu().numpy()) if return_ids: ids = batch.get("id_1", [""] * len(y_t)) eyes = batch.get("eye_id_1", [""] * len(y_t)) if torch.is_tensor(ids): ids = ids.tolist() if torch.is_tensor(eyes): eyes = eyes.tolist() id_c.extend([f"{p}{e}" for p, e in zip(ids, eyes)]) if not y_c: z = np.zeros((0, 0), dtype=np.float32) if return_ids: return np.array([], dtype=np.int64), z, z, z, np.array([], dtype=object) return np.array([], dtype=np.int64), z, z, z y = np.concatenate(y_c) pf, pi, pm = (np.concatenate(c, axis=0) for c in (pf_c, pi_c, pm_c)) if return_ids: return y, pf, pi, pm, np.array(id_c, dtype=object) return y, pf, pi, pm # --------------------------------------------------------------------------- # MonoTowerHT training / inference helpers # --------------------------------------------------------------------------- def train_mono_epoch( model: MonoTowerHT, loader: DataLoader, opt, device: torch.device, ) -> tuple[float, float]: """One training epoch for MonoTowerHT. Uses the tower's ``embed_batch`` (TowerBase API) to extract the embedding, then classifies directly with the head. Loss is cross-entropy on the head output only — no aux head. Returns ------- (mean_loss, accuracy) """ model.train() total_loss = total_correct = total_n = 0 for batch in loader: y = _to_label_tensor(batch.get("label_1"), device) if y.numel() == 0: continue [z] = model.tower.embed_batch(batch, device=device) logits = model.classify(z) loss = F.cross_entropy(logits, y) opt.zero_grad() loss.backward() opt.step() total_loss += loss.item() * len(y) total_correct += int((logits.argmax(1) == y).sum()) total_n += len(y) mean_loss = total_loss / total_n if total_n else float("nan") accuracy = total_correct / total_n if total_n else float("nan") return mean_loss, accuracy def collect_probs_mono( model: MonoTowerHT, loader: DataLoader, device: torch.device, ) -> tuple[np.ndarray, np.ndarray]: """Collect predictions from a MonoTowerHT. Returns ------- y_true : int64 array [N] probs : float32 array [N, num_classes] — softmax of fused head """ model.eval() y_all, p_all = [], [] with torch.no_grad(): for batch in loader: y = _to_label_tensor(batch.get("label_1"), device) if y.numel() == 0: continue [z] = model.tower.embed_batch(batch, device=device) logits = model.classify(z) y_all.append(y.cpu().numpy()) p_all.append(F.softmax(logits, dim=1).cpu().numpy()) if not y_all: return np.zeros(0, dtype=np.int64), np.zeros((0, 0), dtype=np.float32) return np.concatenate(y_all), np.concatenate(p_all, axis=0) # --------------------------------------------------------------------------- # NTowerHT training / inference helpers # --------------------------------------------------------------------------- def train_ntower_epoch( model: NTowerHT, loader: DataLoader, opt, device: torch.device, *, batch_key_map: dict[str, str], phase: str, bcd_prob: float = 0.5, tower_loss_mode: str = "bcd", ) -> tuple[float, float]: """One training epoch for NTowerHT. Parameters ---------- batch_key_map : {tower_name: batch_key} — maps each tower slot to the key in the DataLoader batch dict that carries its input tensor. Example: {'od': 'image_1', 'os': 'image_2', 'cd': 'matrix_1'} phase : 'tower_warmup' | 'fused_warmup' | 'main' bcd_prob : probability of using a random aux head instead of the fused head in 'main' phase (BCD training). tower_loss_mode : 'bcd' | 'all' — 'all' adds all aux + fused losses each step. """ from random import choice model.train() model.bridge.set_phase(phase) # Mirror _set_single_phase tower freezing: # tower_warmup → towers trainable; fused_warmup → towers frozen; main → trainable towers_trainable = phase != "fused_warmup" for t in model.towers.values(): for p in t.parameters(): p.requires_grad_(towers_trainable) tower_names = list(model.towers.keys()) total_loss = total_correct = total_n = 0 for batch in loader: tensors = {name: batch.get(key) for name, key in batch_key_map.items()} if not all(torch.is_tensor(t) for t in tensors.values()): continue y = _to_label_tensor(batch.get("label_1"), device) if y.numel() == 0: continue embeddings = {name: model.towers[name](t.to(device)) for name, t in tensors.items()} if phase == "tower_warmup": aux_logits = [model.bridge.aux_heads[i](embeddings[name]) for i, name in enumerate(tower_names)] loss = sum(F.cross_entropy(l, y) for l in aux_logits) / len(aux_logits) avg_probs = sum(F.softmax(l, dim=1) for l in aux_logits) / len(aux_logits) logits = avg_probs # for accuracy tracking elif phase == "fused_warmup": logits, _ = model(embeddings) loss = F.cross_entropy(logits, y) else: # main if tower_loss_mode == "all": logits, aux_dict = model(embeddings) loss = F.cross_entropy(logits, y) + sum( F.cross_entropy(a, y) for a in aux_dict.values() ) elif random() < bcd_prob: # BCD: pick one random aux head name = choice(tower_names) idx = tower_names.index(name) logits = model.bridge.aux_heads[idx](embeddings[name]) loss = F.cross_entropy(logits, y) else: logits, _ = model(embeddings) loss = F.cross_entropy(logits, y) opt.zero_grad() loss.backward() opt.step() if logits.ndim == 2: total_correct += int((logits.argmax(1) == y).sum()) total_loss += loss.item() * len(y) total_n += len(y) return ( total_loss / total_n if total_n else float("nan"), total_correct / total_n if total_n else float("nan"), ) def collect_probs_ntower( model: NTowerHT, loader: DataLoader, device: torch.device, *, batch_key_map: dict[str, str], ) -> tuple[np.ndarray, np.ndarray]: """Collect fused-head predictions from NTowerHT. Parameters ---------- batch_key_map : same mapping used during training. Returns ------- y_true : int64 array [N] probs : float32 array [N, num_classes] — softmax of fused head """ model.eval() y_all, p_all = [], [] with torch.no_grad(): for batch in loader: tensors = {name: batch.get(key) for name, key in batch_key_map.items()} if not all(torch.is_tensor(t) for t in tensors.values()): continue y = _to_label_tensor(batch.get("label_1"), device) if y.numel() == 0: continue embeddings = {name: model.towers[name](t.to(device)) for name, t in tensors.items()} logits, _ = model(embeddings) y_all.append(y.cpu().numpy()) p_all.append(F.softmax(logits, dim=1).cpu().numpy()) if not y_all: return np.zeros(0, dtype=np.int64), np.zeros((0, 0), dtype=np.float32) return np.concatenate(y_all), np.concatenate(p_all, axis=0) # --------------------------------------------------------------------------- # V2ModeComparisonOps — backward-compat namespace # --------------------------------------------------------------------------- class V2ModeComparisonOps: """Namespace kept for backward-compatibility imports.""" _set_requires_grad = staticmethod(_set_requires_grad) _set_single_phase = staticmethod(_set_single_phase) _set_bilateral_phase = staticmethod(_set_bilateral_phase) train_single_epoch = staticmethod(train_single_epoch) train_bilateral_epoch = staticmethod(train_bilateral_epoch) collect_probs_classic = staticmethod(collect_probs_classic) collect_probs_ensemble = staticmethod(collect_probs_ensemble) collect_probs_bilateral = staticmethod(collect_probs_bilateral) @staticmethod def _to_label_tensor(labels, device: torch.device) -> torch.Tensor: return _to_label_tensor(labels, device)