163 lines
5.0 KiB
Python
163 lines
5.0 KiB
Python
"""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
|