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.
This commit is contained in:
rpotter6298
2026-07-03 08:51:44 +02:00
parent 3d954a4606
commit 708fbc70ce
52 changed files with 2223 additions and 218 deletions
+72 -7
View File
@@ -71,6 +71,13 @@ class ImageEncoder(TowerBase):
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.
@@ -89,6 +96,10 @@ class ImageEncoder(TowerBase):
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,
):
@@ -98,17 +109,37 @@ class ImageEncoder(TowerBase):
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)
_, self._post_eval_tf = build_split_transforms(backbone, augment=False)
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)
self.eval_transform = build_eval_transform(backbone)
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"):
@@ -134,16 +165,32 @@ class ImageEncoder(TowerBase):
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.image_data.load_image(*ids))
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.image_data.load_image(*ids)
img = self._load_image(*ids)
t = self.transform if self.training else self.eval_transform
return t(img)
@@ -154,6 +201,24 @@ class ImageEncoder(TowerBase):
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)
@@ -200,7 +265,7 @@ class ImageEncoder(TowerBase):
key = (pid, eye)
if key in self._tensor_cache or key in seen:
continue
self._tensor_cache[key] = self._precache_tf(self.image_data.load_image(pid, eye))
self._tensor_cache[key] = self._precache_tf(self._load_image(pid, eye))
seen.add(key)
return len(self._tensor_cache)