Files
hypertower/v4/classes/accessory/transforms.py
T
rpotter6298 280060db82 Add new regression and ensemble experiment configurations for V2-M and OrthoBridge
- Introduced multiple regression experiment configurations targeting vf_md, including:
  - cd_solo_reg_set.json: CD tower only regression setup.
  - img_solo_reg_set.json: Image tower only regression setup.
  - reg_head_epoch_sweep.json: Baseline regression sweeps at different epochs (50, 75, 100).
  - reg_head_set.json: Various regression setups including baseline and OrthoBridge configurations.
  - single_eye_reg.json: Single-eye regression setup for worst-eye aggregation analysis.

- Added ensemble configurations for OrthoBridge with different inner bridges:
  - ortho_alts_ensemble.json: Ensemble tests with ConcatBridge, PairwiseAdditiveBridge, and GatedAdditiveBridge.
  - ortho_alts_tritower.json: Tritower tests with the same inner bridges.

- Created V2-M specific configurations:
  - baseline_reg_nt50.json: Regression baseline with V2-M backbone.
  - geom_vec_gt.json and geom_vec_unet.json: Geometry vector injection experiments with V2-M.
  - single_l1_bridges.json: Single-eye ensemble experiments with various bridge types.
  - tritower_geom_gt.json: Tritower setup with GT contour-rasterized masks.

- Promoted existing experiments to higher repetitions for robustness.
2026-06-11 15:08:20 +02:00

115 lines
4.8 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) -> ImageTransformConfig:
"""Build an ImageTransformConfig using the backbone's default normalisation stats."""
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.
return ImageTransformConfig(crop_size=224, mean=IMAGENET_MEAN,
std=IMAGENET_STD, augment=augment)
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)
crop = 299 if key == "inception_v3" else 224
return ImageTransformConfig(crop_size=crop, mean=mean, std=std, augment=augment)
def build_backbone_transform(backbone_name: str, augment: bool = True) -> transforms.Compose:
return backbone_transform_config(backbone_name, augment=augment).build()
def build_eval_transform(backbone_name: str) -> transforms.Compose:
"""Deterministic eval transform — no augmentation, backbone-matched normalisation."""
return build_backbone_transform(backbone_name, augment=False)
def build_split_transforms(
backbone_name: str, augment: bool = True
) -> 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)
return cfg.build_precache(), cfg.build_postcache()