Files
hypertower/v4/classes/towers/image_tower.py
T
rpotter6298 708fbc70ce Add analysis scripts and experiment configurations for bridge attention and sensitivity studies
- 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.
2026-07-03 08:51:44 +02:00

292 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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