pre-refactor 041426

This commit is contained in:
rpotter6298
2026-04-14 19:42:16 +02:00
parent eb9eafe715
commit 13290575d5
75 changed files with 8900 additions and 77 deletions
+222 -6
View File
@@ -11,7 +11,7 @@ from torch import nn
from torch.utils.data import DataLoader
from v3.classes.bridges import Bridge
from v3.classes.towers import ImageTower, ClinicalTower
from v3.classes.towers import ImageTower, ClinicalTower, SiameseImageTower
# ---------------------------------------------------------------------------
@@ -35,18 +35,24 @@ class SingleEyeHT(nn.Module):
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 = ImageTower(
backbone=backbone,
freeze_ratio=freeze_ratio,
augment=augment,
use_se=False,
use_se=se_img_tower,
)
self.cd_tower = ClinicalTower(
clinical_data=clinical_data,
hidden_dim=cd_hidden_dim,
use_se=False,
dropout=cd_dropout,
use_se=se_cd_tower,
)
self.bridge = Bridge(
img_dim=self.img_tower.out_dim,
@@ -54,13 +60,20 @@ class SingleEyeHT(nn.Module):
num_classes=num_classes,
fusion_dim=fusion_dim,
mode=bridge_mode,
use_se=False,
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 embedding (fusion_dim) without applying the classifier head."""
img_feats = self.img_tower(x)
md_feats = self.cd_tower(meta)
return self.bridge.encode(img_feats, md_feats)
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)
@@ -158,6 +171,59 @@ class BilateralHT(nn.Module):
return out_f
class SiameseHT(nn.Module):
"""
Bilateral model using a shared-weight SiameseImageTower (mean+delta).
Both eyes pass through the same backbone; features are combined as
cat([mean(f_od, f_os), f_od - f_os]) giving the model both a shared
bilateral representation and an asymmetry signal.
Uses the same bilateral loader and training loop as BilateralHT.
The bridge operates in image_only mode (no clinical data in phase 4).
"""
def __init__(
self,
*,
backbone: str,
freeze_ratio: float,
augment: bool,
num_classes: int,
fusion_dim: int = 256,
bridge_mode: str = "image_only",
):
super().__init__()
self.img_tower = SiameseImageTower(
backbone=backbone,
freeze_ratio=freeze_ratio,
augment=augment,
use_se=False,
)
img_dim = self.img_tower.out_dim # 2 * backbone_out_dim
self.bridge = Bridge(
img_dim=img_dim,
meta_dim=1, # dummy — not used in image_only mode
num_classes=num_classes,
fusion_dim=fusion_dim,
mode="image_only",
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:
feats = self.encode(x_od, x_os)
out_f, _, _ = self.bridge(feats, None)
return out_f
class FusedEnsembleHT(nn.Module):
"""
SingleEyeHT base with a per-eye attention scorer for bilateral fusion.
@@ -190,6 +256,10 @@ class FusedEnsembleHT(nn.Module):
# Learns the GC-direction in logit space from bilateral labels.
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,
@@ -205,6 +275,78 @@ class FusedEnsembleHT(nn.Module):
return alpha[:, 0:1] * logit_od + alpha[:, 1:2] * logit_os # [B, C]
class LogitMLPEnsembleHT(nn.Module):
"""
MLP head trained on concatenated per-eye logits.
Both eyes pass through the frozen base independently, producing per-eye
logit vectors. These are concatenated and fed through a small MLP:
cat([logit_od, logit_os]) [B, 2C]
→ Linear(2C, hidden) → ReLU → Dropout → Linear(hidden, C)
Permutation-variant by design: the model can learn left/right asymmetries
directly from the concatenated pair, at the cost of needing a consistent
OD-first input ordering.
"""
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:
logit_od = self.base(x_od, meta_od)
logit_os = self.base(x_os, meta_os)
return self.head(torch.cat([logit_od, logit_os], dim=1))
class EmbeddingMLPEnsembleHT(nn.Module):
"""
MLP head trained on concatenated per-eye z_fused embeddings.
Both eyes pass through the frozen base independently, and their bridge
embeddings (pre-classifier, shape [B, fusion_dim]) are concatenated and
fed through an MLP:
cat([z_od, z_os]) [B, 2 * fusion_dim]
→ Linear(2*fusion_dim, hidden) → ReLU → Dropout → Linear(hidden, C)
Richer than logit-level: the head sees pre-softmax feature vectors rather
than the compressed C-dimensional output, giving it more signal to work
with when fusion_dim >> C.
"""
def __init__(self, base: SingleEyeHT, num_classes: int, hidden: int = 256):
super().__init__()
self.base = base
fusion_dim = base.bridge.W_img.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:
z_od = self.base.encode(x_od, meta_od)
z_os = self.base.encode(x_os, meta_os)
return self.head(torch.cat([z_od, z_os], dim=1))
# ---------------------------------------------------------------------------
# Phase control
# ---------------------------------------------------------------------------
@@ -429,15 +571,89 @@ def train_bilateral_epoch(
)
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]:
"""Train one epoch of SiameseHT on bilateral (patient-level) samples."""
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":
logits_i = model.aux_img(feats)
out_f, _, _ = model.bridge(feats, None)
loss = F.cross_entropy(out_f, y) + F.cross_entropy(logits_i, y)
logits = out_f
elif random() < bcd_prob:
logits = model.aux_img(feats)
loss = F.cross_entropy(logits, y)
else:
logits, _, _ = model.bridge(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 collect_probs_siamese(
model: SiameseHT,
loader: DataLoader,
device: torch.device,
) -> tuple[np.ndarray, np.ndarray]:
"""Patient-level probs from a bilateral loader using SiameseHT."""
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
x1 = x1.to(device); x2 = x2.to(device)
feats = model.encode(x1, x2)
logits, _, _ = model.bridge(feats, None)
probs = torch.softmax(logits, dim=1)
y_c.append(np.array(y) if not torch.is_tensor(y) else y.cpu().numpy())
p_c.append(probs.cpu().numpy())
return (
np.concatenate(y_c, axis=0),
np.concatenate(p_c, axis=0),
)
def train_fusion_epoch(
model: FusedEnsembleHT,
model, # FusedEnsembleHT | LogitMLPEnsembleHT | EmbeddingMLPEnsembleHT
loader: DataLoader,
opt,
device: torch.device,
) -> tuple[float, float]:
"""Train only the fusion head; the base SingleEyeHT is frozen in eval mode."""
model.base.eval()
model.eye_scorer.train()
model.head.train()
total_loss = total_correct = total_n = 0
for batch in loader:
x1 = batch.get("image_1"); m1 = batch.get("matrix_1")