280 lines
9.4 KiB
Python
280 lines
9.4 KiB
Python
from __future__ import annotations
|
||
|
||
import math
|
||
from typing import Optional
|
||
|
||
import torch
|
||
from torch import nn
|
||
from torchvision import transforms
|
||
|
||
from classes.v2.backbones import BACKBONES, list_names, load_backbone_weights
|
||
from classes.v2.SE_attention import SEBlock
|
||
from classes.v2.data_bundle import DataBundle
|
||
|
||
|
||
def build_backbone(name: str, freeze_ratio: float = 0.0, augment: bool = True):
|
||
"""
|
||
Operational builder:
|
||
- instantiate with DEFAULT weights
|
||
- strip classifier → features
|
||
- apply ratio-based freezing over coarse blocks
|
||
- return (model, out_dim, transform)
|
||
"""
|
||
key = (name or "").lower()
|
||
if key not in BACKBONES:
|
||
raise ValueError(f"Unsupported backbone '{name}'. Valid options: {list_names()}")
|
||
|
||
spec = BACKBONES[key]
|
||
m = spec.ctor(weights=spec.weights_default)
|
||
out_dim, m = spec.strip(m)
|
||
load_backbone_weights(key, m)
|
||
|
||
# transforms: use the weights’ mean/std, but keep your augmentation pipeline
|
||
mean = getattr(spec.weights_default, "meta", {}).get("mean", (0.485, 0.456, 0.406))
|
||
std = getattr(spec.weights_default, "meta", {}).get("std", (0.229, 0.224, 0.225))
|
||
crop = 299 if key == "inception_v3" else 224
|
||
|
||
if augment:
|
||
transform = transforms.Compose(
|
||
[
|
||
transforms.Resize(256),
|
||
transforms.CenterCrop(crop),
|
||
transforms.RandomHorizontalFlip(),
|
||
transforms.RandomVerticalFlip(),
|
||
transforms.RandomRotation(15),
|
||
transforms.ColorJitter(0.1, 0.1, 0.1, 0.05),
|
||
transforms.ToTensor(),
|
||
transforms.Normalize(mean=mean, std=std),
|
||
]
|
||
)
|
||
else:
|
||
transform = transforms.Compose(
|
||
[
|
||
transforms.Resize(256),
|
||
transforms.CenterCrop(crop),
|
||
transforms.ToTensor(),
|
||
transforms.Normalize(mean=mean, std=std),
|
||
]
|
||
)
|
||
|
||
# ratio-based freezing: freeze earliest floor(N * freeze_ratio) blocks
|
||
fr = max(0.0, min(1.0, float(freeze_ratio)))
|
||
blocks = spec.blocks(m)
|
||
n = len(blocks)
|
||
freeze_n = int(math.floor(n * fr))
|
||
for b in blocks[:freeze_n]:
|
||
for p in b.parameters():
|
||
p.requires_grad = False
|
||
|
||
return m, out_dim, transform
|
||
|
||
|
||
class ImageTower(nn.Module):
|
||
"""
|
||
Vision backbone → pooled features.
|
||
- backbone: one of list_names() (default 'efficientnet_b0')
|
||
- always DEFAULT torchvision weights
|
||
- freeze_ratio ∈ [0,1] freezes earliest floor(N*freeze_ratio) blocks
|
||
- returns [N, out_dim] features from backbone forward
|
||
"""
|
||
|
||
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
|
||
# Keep ordered blocks for dynamic freezing/thawing
|
||
key = (self._name or "").lower()
|
||
self._spec = BACKBONES[key]
|
||
self._blocks = self._spec.blocks(self.backbone)
|
||
# Optional tower-level SE over the final feature vector
|
||
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)
|
||
# sanity: pooled features, not logits
|
||
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:
|
||
if geometry.dim() == 1:
|
||
geom = geometry.unsqueeze(0)
|
||
else:
|
||
geom = 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):
|
||
"""Dynamically freeze earliest floor(N*ratio) backbone blocks."""
|
||
r = max(0.0, min(1.0, float(ratio)))
|
||
n = len(self._blocks)
|
||
freeze_n = int(math.floor(n * r))
|
||
# Unfreeze all first
|
||
for b in self._blocks:
|
||
for p in b.parameters():
|
||
p.requires_grad = True
|
||
# Freeze earliest blocks
|
||
for b in self._blocks[:freeze_n]:
|
||
for p in b.parameters():
|
||
p.requires_grad = False
|
||
|
||
|
||
class SiameseImageTower(nn.Module):
|
||
"""
|
||
Shared-weight bilateral image tower.
|
||
|
||
Runs OD and OS images through a single shared backbone, then returns
|
||
cat([f_mean, f_delta]) where:
|
||
f_mean = (f_od + f_os) / 2 -- shared bilateral representation
|
||
f_delta = f_od - f_os -- asymmetry, signed OD-relative
|
||
|
||
out_dim = 2 * backbone_out_dim
|
||
|
||
When x_os is None (single-eye fallback):
|
||
f_mean = f_od
|
||
f_delta = zeros
|
||
so the module degrades gracefully when only one eye is available.
|
||
|
||
The shared backbone means both eyes contribute to every gradient update,
|
||
effectively doubling the training signal for the visual pathway without
|
||
doubling parameters.
|
||
"""
|
||
|
||
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 = ImageTower(
|
||
backbone=backbone,
|
||
freeze_ratio=freeze_ratio,
|
||
use_se=use_se,
|
||
se_reduction=se_reduction,
|
||
se_pre_norm=se_pre_norm,
|
||
augment=augment,
|
||
geometry_dim=0,
|
||
)
|
||
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:
|
||
f_mean = f_od
|
||
f_delta = torch.zeros_like(f_od)
|
||
else:
|
||
f_os = self._tower(x_os)
|
||
f_mean = (f_od + f_os) * 0.5
|
||
f_delta = f_od - f_os
|
||
return torch.cat([f_mean, f_delta], dim=1)
|
||
|
||
def set_freeze_ratio(self, ratio: float) -> None:
|
||
"""Delegates to the shared inner tower."""
|
||
self._tower.set_freeze_ratio(ratio)
|
||
|
||
|
||
class MDTower(nn.Module):
|
||
"""MLP over DataBundle.vectorize_row outputs (convert 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)))
|
||
# Unfreeze all
|
||
for p in self.block0.parameters():
|
||
p.requires_grad = True
|
||
for p in self.block1.parameters():
|
||
p.requires_grad = True
|
||
# Freeze earliest blocks based on ratio threshold
|
||
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
|