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:
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user