253 lines
8.6 KiB
Python
253 lines
8.6 KiB
Python
"""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,
|
||
)
|