pre-refactor 041426
This commit is contained in:
+222
-6
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user