Files
hypertower/v4/classes/accessory/transforms.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

151 lines
5.7 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.
"""transforms — image transform utilities for v4 towers."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Tuple
from torchvision import transforms
from v4.classes.accessory.backbones import BACKBONES, _is_timm_backbone
IMAGENET_MEAN: Tuple[float, float, float] = (0.485, 0.456, 0.406)
IMAGENET_STD: Tuple[float, float, float] = (0.229, 0.224, 0.225)
@dataclass
class ImageTransformConfig:
crop_size: int = 224
resize_size: int = 256
mean: Tuple[float, float, float] = IMAGENET_MEAN
std: Tuple[float, float, float] = IMAGENET_STD
augment: bool = True
rotation_deg: int = 15
color_jitter: Tuple[float, float, float, float] = (0.1, 0.1, 0.1, 0.05)
hflip: bool = True
vflip: bool = True
def build(self) -> transforms.Compose:
ops = [
transforms.Resize(self.resize_size),
transforms.CenterCrop(self.crop_size),
]
if self.augment:
if self.hflip:
ops.append(transforms.RandomHorizontalFlip())
if self.vflip:
ops.append(transforms.RandomVerticalFlip())
if self.rotation_deg:
ops.append(transforms.RandomRotation(self.rotation_deg))
if self.color_jitter:
ops.append(transforms.ColorJitter(*self.color_jitter))
ops += [
transforms.ToTensor(),
transforms.Normalize(mean=self.mean, std=self.std),
]
return transforms.Compose(ops)
def build_precache(self) -> transforms.Compose:
"""Deterministic prefix: PIL → resized CHW float32 in [0, 1].
Output is suitable for caching; per-batch ``build_postcache`` finishes
the pipeline (augment + normalize) on tensors.
"""
return transforms.Compose([
transforms.Resize(self.resize_size),
transforms.CenterCrop(self.crop_size),
transforms.ToTensor(),
])
def build_postcache(self) -> transforms.Compose:
"""Per-batch tail run on cached float32 [0, 1] CHW tensors.
Augmentations operate on tensors (torchvision v1 supports this for
Flip/Rotation/ColorJitter on tensor input). Normalize is applied last.
"""
ops = []
if self.augment:
if self.hflip:
ops.append(transforms.RandomHorizontalFlip())
if self.vflip:
ops.append(transforms.RandomVerticalFlip())
if self.rotation_deg:
ops.append(transforms.RandomRotation(self.rotation_deg))
if self.color_jitter:
ops.append(transforms.ColorJitter(*self.color_jitter))
ops.append(transforms.Normalize(mean=self.mean, std=self.std))
return transforms.Compose(ops)
def backbone_transform_config(
backbone_name: str,
augment: bool = True,
crop_size: int | None = None,
resize_size: int | None = None,
) -> ImageTransformConfig:
"""Build an ImageTransformConfig using the backbone's default normalisation stats.
crop_size / resize_size override the backbone's default input resolution. When
crop_size is overridden but resize_size is not, resize_size is scaled
proportionally (8/7 ratio, matching the standard 224 → 256 pattern).
"""
key = (backbone_name or "").lower()
if _is_timm_backbone(key):
# ConvNeXt-V2 and other timm models we currently expose are all
# pretrained with standard ImageNet stats at 224×224.
mean, std = IMAGENET_MEAN, IMAGENET_STD
default_crop = 224
else:
if key not in BACKBONES:
raise ValueError(f"Unknown backbone '{backbone_name}'.")
spec = BACKBONES[key]
mean = getattr(spec.weights_default, "meta", {}).get("mean", IMAGENET_MEAN)
std = getattr(spec.weights_default, "meta", {}).get("std", IMAGENET_STD)
default_crop = 299 if key == "inception_v3" else 224
crop = crop_size if crop_size is not None else default_crop
resize = resize_size if resize_size is not None else round(crop * 8 / 7)
return ImageTransformConfig(crop_size=crop, resize_size=resize,
mean=mean, std=std, augment=augment)
def build_backbone_transform(
backbone_name: str,
augment: bool = True,
crop_size: int | None = None,
resize_size: int | None = None,
) -> transforms.Compose:
return backbone_transform_config(
backbone_name, augment=augment,
crop_size=crop_size, resize_size=resize_size,
).build()
def build_eval_transform(
backbone_name: str,
crop_size: int | None = None,
resize_size: int | None = None,
) -> transforms.Compose:
"""Deterministic eval transform — no augmentation, backbone-matched normalisation."""
return build_backbone_transform(
backbone_name, augment=False,
crop_size=crop_size, resize_size=resize_size,
)
def build_split_transforms(
backbone_name: str,
augment: bool = True,
crop_size: int | None = None,
resize_size: int | None = None,
) -> tuple[transforms.Compose, transforms.Compose]:
"""Return (precache, postcache) transform pair for tensor-cached image towers.
precache : PIL → CHW float32 in [0, 1] (deterministic, run once at fill)
postcache : tensor → augmented + normalized tensor (run per batch)
"""
cfg = backbone_transform_config(
backbone_name, augment=augment,
crop_size=crop_size, resize_size=resize_size,
)
return cfg.build_precache(), cfg.build_postcache()