708fbc70ce
- Introduced `bridge_attention_ceiling_check.py` for variance decomposition analysis on bridge attention configurations. - Added `bridge_attention_readout.py` to perform per-tower gate and contribution readouts, including AUC sanity checks. - Created multiple JSON configuration files for backbone replication experiments, including anonymous CV variants and basic backbones. - Implemented sensitivity experiments to evaluate the impact of axial length inclusion and EfficientNetV2-M performance at higher resolutions. - Added a memory probe script to assess GPU memory usage during training with EfficientNetV2-M.
292 lines
12 KiB
Python
292 lines
12 KiB
Python
"""image_tower — ImageEncoder for v4.
|
||
|
||
Self-contained: no v3 dependencies.
|
||
Inherits get_sample dispatch from TowerBase.
|
||
|
||
Geometry injection (EPC supply)
|
||
--------------------------------
|
||
When geometry_source is set, ImageEncoder asks the image_data view for a loader
|
||
via image_data.build_geometry_loader(source, **kwargs). The view is responsible
|
||
for understanding what that source means for its specific domain (fundus contours,
|
||
U-Net segmentations, cat ear landmarks, etc.).
|
||
|
||
During early_pass the loader pre-computes all per-entity geometry vectors and
|
||
publishes them to the EarlyPassContext under the key "geometry_vectors"
|
||
({(entity_id...): np.ndarray of length geom_dim}). ClinicalEncoder (or any
|
||
other tower with epc_requests: ["geometry_vectors"]) can then consume them.
|
||
|
||
The tower reads feature_dim and feature_names from the loader instance, so it
|
||
can log geometry info without knowing anything about CDR, disc masks, or other
|
||
domain-specific concepts.
|
||
|
||
Config example:
|
||
{
|
||
"name": "img",
|
||
"module": "v4.classes.towers.image_tower",
|
||
"class": "ImageEncoder",
|
||
"data_source": "image",
|
||
"epc_supplies": ["geometry_vectors"],
|
||
"args": {
|
||
"backbone": "refugelike",
|
||
"augment": true,
|
||
"geometry_source": "gt",
|
||
"contour_dir": "Papila/ExpertsSegmentations/Contours"
|
||
}
|
||
}
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import math
|
||
import sys
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
import torch
|
||
from torch import nn
|
||
|
||
_REPO_ROOT = Path(__file__).resolve().parents[3]
|
||
if str(_REPO_ROOT) not in sys.path:
|
||
sys.path.insert(0, str(_REPO_ROOT))
|
||
|
||
from v4.classes.towerbase import TowerBase
|
||
from v4.classes.accessory.backbones import build_backbone
|
||
from v4.classes.accessory.se_block import SEBlock
|
||
from v4.classes.accessory.transforms import (
|
||
build_backbone_transform, build_eval_transform, build_split_transforms,
|
||
)
|
||
|
||
|
||
class ImageEncoder(TowerBase):
|
||
"""Vision backbone → pooled feature vector.
|
||
|
||
image_data : ImageDataView — provides load_image(*ids) and side_map.
|
||
Must implement build_geometry_loader(source, **kwargs)
|
||
if geometry_source is set.
|
||
backbone : backbone key (see accessory/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 train transform
|
||
cache_transformed : if True, cache resized + ToTensor'd float32 [0, 1] CHW
|
||
tensors per fold. Per-batch cost drops to augment +
|
||
Normalize on tensors only (no PIL, no Resize, no decode).
|
||
Memory: ~3 × crop_size² × 4B per cached image.
|
||
Cache is rebuilt at the start of every fold via early_pass.
|
||
crop_source : if set, crop each input image to a square disc-region
|
||
bbox before the standard transform pipeline. Values:
|
||
"gt" (use GT contour file) | "unet" (use U-Net mask)
|
||
| None (disabled, full-image pipeline).
|
||
crop_kwargs : dict forwarded to image_data.build_disc_bbox_loader().
|
||
Common keys: margin (default 2.5), expert (GT only),
|
||
weights_path / finetune_epochs (U-Net only).
|
||
geometry_source : source key passed to image_data.build_geometry_loader()
|
||
(e.g. "gt", "unet"). None = geometry disabled.
|
||
**geom_kwargs : forwarded verbatim to build_geometry_loader() — e.g.
|
||
contour_dir="Papila/ExpertsSegmentations/Contours"
|
||
"""
|
||
|
||
EPC_GEOMETRY_KEY = "geometry_vectors"
|
||
|
||
def __init__(
|
||
self,
|
||
image_data,
|
||
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,
|
||
cache_transformed: bool = False,
|
||
crop_size: int | None = None,
|
||
resize_size: int | None = None,
|
||
crop_source: str | None = None,
|
||
crop_kwargs: dict | None = None,
|
||
geometry_source: str | None = None,
|
||
**geom_kwargs: Any,
|
||
):
|
||
super().__init__()
|
||
self.image_data = image_data
|
||
self._name = backbone
|
||
self.backbone, self._base_dim, self._blocks = build_backbone(backbone, freeze_ratio)
|
||
|
||
self._cache_transformed = cache_transformed
|
||
tf_kw = dict(crop_size=crop_size, resize_size=resize_size)
|
||
if cache_transformed:
|
||
self._precache_tf, self._post_train_tf = build_split_transforms(
|
||
backbone, augment=augment, **tf_kw)
|
||
_, self._post_eval_tf = build_split_transforms(
|
||
backbone, augment=False, **tf_kw)
|
||
self._tensor_cache: dict[tuple, torch.Tensor] = {}
|
||
else:
|
||
self.transform = build_backbone_transform(backbone, augment=augment, **tf_kw)
|
||
self.eval_transform = build_eval_transform(backbone, **tf_kw)
|
||
|
||
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
|
||
|
||
self._bbox_loader = None
|
||
if crop_source is not None:
|
||
if not hasattr(image_data, "build_disc_bbox_loader"):
|
||
raise TypeError(
|
||
f"ImageEncoder crop_source={crop_source!r} requires "
|
||
f"image_data to implement build_disc_bbox_loader(), "
|
||
f"but {type(image_data).__name__} does not."
|
||
)
|
||
self._bbox_loader = image_data.build_disc_bbox_loader(
|
||
crop_source, **(crop_kwargs or {}),
|
||
)
|
||
print(
|
||
f"[ImageEncoder] crop_source={crop_source!r} "
|
||
f"kwargs={crop_kwargs or {}}",
|
||
flush=True,
|
||
)
|
||
|
||
self._geom_loader = None
|
||
if geometry_source is not None:
|
||
if not hasattr(image_data, "build_geometry_loader"):
|
||
raise TypeError(
|
||
f"ImageEncoder geometry_source={geometry_source!r} requires "
|
||
f"image_data to implement build_geometry_loader(), "
|
||
f"but {type(image_data).__name__} does not."
|
||
)
|
||
self._geom_loader = image_data.build_geometry_loader(geometry_source, **geom_kwargs)
|
||
print(
|
||
f"[ImageEncoder] geometry_source={geometry_source!r} "
|
||
f"features={self._geom_loader.feature_names}",
|
||
flush=True,
|
||
)
|
||
|
||
# ── TowerBase interface ──────────────────────────────────────────────────
|
||
|
||
@property
|
||
def out_dim(self) -> int:
|
||
return self._base_dim
|
||
|
||
@property
|
||
def _side_map(self) -> dict[str, str]:
|
||
return self.image_data.side_map
|
||
|
||
def _load_image(self, *ids):
|
||
"""Load image, optionally cropped to the disc-region bbox."""
|
||
pil = self.image_data.load_image(*ids)
|
||
if self._bbox_loader is None:
|
||
return pil
|
||
bbox = self._bbox_loader.bbox_for(*ids[:2])
|
||
if bbox is None:
|
||
return pil
|
||
w, h = pil.size
|
||
x0, y0, x1, y1 = bbox
|
||
x0, y0 = max(0, x0), max(0, y0)
|
||
x1, y1 = min(w, x1), min(h, y1)
|
||
if x1 <= x0 or y1 <= y0:
|
||
return pil
|
||
return pil.crop((x0, y0, x1, y1))
|
||
|
||
def _get(self, *ids) -> torch.Tensor:
|
||
if self._cache_transformed:
|
||
key = tuple(ids)
|
||
cached = self._tensor_cache.get(key)
|
||
if cached is None:
|
||
cached = self._precache_tf(self._load_image(*ids))
|
||
self._tensor_cache[key] = cached
|
||
tail = self._post_train_tf if self.training else self._post_eval_tf
|
||
return tail(cached)
|
||
img = self._load_image(*ids)
|
||
t = self.transform if self.training else self.eval_transform
|
||
return t(img)
|
||
|
||
# ── EPC early_pass ───────────────────────────────────────────────────────
|
||
|
||
def early_pass(self, context) -> None:
|
||
"""Per-fold setup: warm tensor cache (if enabled), publish geometry vectors."""
|
||
data = context.require("data")
|
||
split = context.require("split")
|
||
|
||
# Disc-region bbox precomputation (must run before any image load/cache).
|
||
if self._bbox_loader is not None:
|
||
train_samples = data.collect_samples(split.train)
|
||
all_samples = train_samples + data.collect_samples(split.val)
|
||
if split.test is not None:
|
||
all_samples += data.collect_samples(split.test)
|
||
if hasattr(self._bbox_loader, "reset_cache"):
|
||
self._bbox_loader.reset_cache()
|
||
if hasattr(self._bbox_loader, "reset_weights"):
|
||
self._bbox_loader.reset_weights()
|
||
if hasattr(self._bbox_loader, "finetune"):
|
||
self._bbox_loader.finetune(train_samples)
|
||
self._bbox_loader.precompute(all_samples)
|
||
print(
|
||
f"[ImageEncoder] precomputed disc bboxes for {len(all_samples)} samples",
|
||
flush=True,
|
||
)
|
||
|
||
if self._cache_transformed:
|
||
self._tensor_cache.clear()
|
||
n = self._warm_tensor_cache(data, split)
|
||
print(
|
||
f"[ImageEncoder] warmed transformed-tensor cache for {n} entries "
|
||
f"({self._name})",
|
||
flush=True,
|
||
)
|
||
|
||
if self._geom_loader is None:
|
||
return
|
||
|
||
train_samples = data.collect_samples(split.train)
|
||
all_samples = train_samples + data.collect_samples(split.val)
|
||
if split.test is not None:
|
||
all_samples += data.collect_samples(split.test)
|
||
|
||
if hasattr(self._geom_loader, "reset_cache"):
|
||
self._geom_loader.reset_cache()
|
||
if hasattr(self._geom_loader, "reset_weights"):
|
||
self._geom_loader.reset_weights()
|
||
if hasattr(self._geom_loader, "finetune"):
|
||
self._geom_loader.finetune(train_samples)
|
||
|
||
self._geom_loader.precompute(all_samples)
|
||
vecs = self._geom_loader.all_vectors()
|
||
context.put(self.EPC_GEOMETRY_KEY, vecs)
|
||
print(
|
||
f"[ImageEncoder] published {len(vecs)} geometry vectors "
|
||
f"(dim={self._geom_loader.feature_dim}) to EPC key '{self.EPC_GEOMETRY_KEY}'",
|
||
flush=True,
|
||
)
|
||
|
||
def _warm_tensor_cache(self, data, split) -> int:
|
||
"""Pre-fill the per-tower tensor cache for all entries in this fold's splits."""
|
||
seen: set[tuple] = set()
|
||
for df in (split.train, split.val, split.test):
|
||
if df is None or len(df) == 0:
|
||
continue
|
||
pc = data.patient_col
|
||
for _, row in df.iterrows():
|
||
pid = int(row[pc])
|
||
eye = str(row.get("eyeID", "OD"))
|
||
key = (pid, eye)
|
||
if key in self._tensor_cache or key in seen:
|
||
continue
|
||
self._tensor_cache[key] = self._precache_tf(self._load_image(pid, eye))
|
||
seen.add(key)
|
||
return len(self._tensor_cache)
|
||
|
||
# ── nn.Module forward ────────────────────────────────────────────────────
|
||
|
||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||
y = self.backbone(x)
|
||
if self.tower_se is not None:
|
||
y, _ = self.tower_se(self.tower_ln(y))
|
||
return y
|
||
|
||
# ── Utilities ────────────────────────────────────────────────────────────
|
||
|
||
def set_freeze_ratio(self, ratio: float) -> None:
|
||
"""Dynamically freeze the earliest floor(N * ratio) backbone blocks."""
|
||
r = max(0.0, min(1.0, float(ratio)))
|
||
n_freeze = 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[:n_freeze]:
|
||
for p in b.parameters():
|
||
p.requires_grad = False
|