moved_repo_first_update
This commit is contained in:
@@ -0,0 +1,315 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Iterable, Optional, Tuple, Union
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from torchvision import transforms
|
||||
|
||||
from classes.backbones import BACKBONES
|
||||
|
||||
|
||||
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:
|
||||
"""
|
||||
Mirrors the hypertower v1 preprocessing:
|
||||
- Resize(256)
|
||||
- CenterCrop(crop)
|
||||
- Optional augmentations (H/V flip, rotation, color jitter)
|
||||
- ToTensor + Normalize(mean/std)
|
||||
"""
|
||||
|
||||
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.extend(
|
||||
[
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=self.mean, std=self.std),
|
||||
]
|
||||
)
|
||||
return transforms.Compose(ops)
|
||||
|
||||
|
||||
def backbone_transform_config(backbone_name: str, augment: bool = True) -> ImageTransformConfig:
|
||||
"""
|
||||
Build a transform config that matches v1 ImageTower/backbone preprocessing.
|
||||
Uses DEFAULT weights mean/std and InceptionV3 crop size when relevant.
|
||||
"""
|
||||
key = (backbone_name or "").lower()
|
||||
if key not in BACKBONES:
|
||||
raise ValueError(f"Unsupported 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_imagenet_transform(augment: bool = True, crop_size: int = 224) -> transforms.Compose:
|
||||
return ImageTransformConfig(crop_size=crop_size, augment=augment).build()
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResizeTransform:
|
||||
size: Union[int, Tuple[int, int]] = 256
|
||||
interpolation: int = Image.BILINEAR
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self._op = transforms.Resize(self.size, interpolation=self.interpolation)
|
||||
|
||||
def __call__(self, image: Image.Image) -> Image.Image:
|
||||
return self._op(image)
|
||||
|
||||
|
||||
@dataclass
|
||||
class CenterCropTransform:
|
||||
size: Union[int, Tuple[int, int]] = 224
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self._op = transforms.CenterCrop(self.size)
|
||||
|
||||
def __call__(self, image: Image.Image) -> Image.Image:
|
||||
return self._op(image)
|
||||
|
||||
|
||||
class UnetMaskProvider:
|
||||
"""
|
||||
Placeholder for a UNet-powered mask provider.
|
||||
This will be replaced once a UNet tower is wired in.
|
||||
"""
|
||||
|
||||
def __call__(self, image: Image.Image, image_path: Optional[str] = None):
|
||||
raise NotImplementedError("UNet mask provider is not wired yet.")
|
||||
|
||||
|
||||
@dataclass
|
||||
class ROICropTransform:
|
||||
"""
|
||||
Crop an image using a binary mask (GT or UNet).
|
||||
Expects a mask of the same spatial size as the image; nonzero pixels are ROI.
|
||||
"""
|
||||
|
||||
mask_source: str = "gt" # "gt" | "unet"
|
||||
mask_provider: Optional[Callable[[Image.Image, Optional[str]], np.ndarray]] = None
|
||||
scale: float = 2.5
|
||||
target_size: Optional[Tuple[int, int]] = (224, 224)
|
||||
fallback_to_original: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.mask_source not in {"gt", "unet"}:
|
||||
raise ValueError(f"mask_source must be 'gt' or 'unet', got '{self.mask_source}'.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
image: Image.Image,
|
||||
mask: Optional[Union[np.ndarray, Image.Image]] = None,
|
||||
image_path: Optional[str] = None,
|
||||
) -> Image.Image:
|
||||
resolved_mask = mask
|
||||
if resolved_mask is None and self.mask_provider is not None:
|
||||
resolved_mask = self.mask_provider(image, image_path)
|
||||
if resolved_mask is None:
|
||||
if self.fallback_to_original:
|
||||
return image
|
||||
raise ValueError("ROI crop requested but no mask provided.")
|
||||
|
||||
mask_arr = (
|
||||
np.asarray(resolved_mask)
|
||||
if not isinstance(resolved_mask, Image.Image)
|
||||
else np.array(resolved_mask)
|
||||
)
|
||||
if mask_arr.ndim == 3:
|
||||
mask_arr = mask_arr[..., 0]
|
||||
mask_arr = mask_arr > 0
|
||||
if not np.any(mask_arr):
|
||||
return image if self.fallback_to_original else image
|
||||
|
||||
ys, xs = np.where(mask_arr)
|
||||
y_min, y_max = ys.min(), ys.max()
|
||||
x_min, x_max = xs.min(), xs.max()
|
||||
cx = (x_min + x_max) / 2.0
|
||||
cy = (y_min + y_max) / 2.0
|
||||
width = (x_max - x_min + 1)
|
||||
height = (y_max - y_min + 1)
|
||||
size = max(width, height) * float(self.scale)
|
||||
|
||||
left = int(round(cx - size / 2))
|
||||
right = int(round(cx + size / 2))
|
||||
upper = int(round(cy - size / 2))
|
||||
lower = int(round(cy + size / 2))
|
||||
|
||||
left = max(0, left)
|
||||
upper = max(0, upper)
|
||||
right = min(image.width, right)
|
||||
lower = min(image.height, lower)
|
||||
crop = image.crop((left, upper, right, lower))
|
||||
if self.target_size is not None:
|
||||
crop = crop.resize(self.target_size, Image.BILINEAR)
|
||||
return crop
|
||||
|
||||
|
||||
@dataclass
|
||||
class JitterBundleTransform:
|
||||
"""
|
||||
Augmentations bundle: flips, rotation, color jitter.
|
||||
"""
|
||||
|
||||
hflip: bool = True
|
||||
vflip: bool = True
|
||||
rotation_deg: int = 15
|
||||
color_jitter: Optional[Tuple[float, float, float, float]] = (0.1, 0.1, 0.1, 0.05)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
ops = []
|
||||
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))
|
||||
self._op = transforms.Compose(ops) if ops else None
|
||||
|
||||
def __call__(self, image: Image.Image) -> Image.Image:
|
||||
if self._op is None:
|
||||
return image
|
||||
return self._op(image)
|
||||
|
||||
|
||||
TRANSFORM_REGISTRY = {
|
||||
"resize": ResizeTransform,
|
||||
"roi_crop": ROICropTransform,
|
||||
"center_crop": CenterCropTransform,
|
||||
"jitter_bundle": JitterBundleTransform,
|
||||
}
|
||||
|
||||
|
||||
def _parse_color_jitter(value: Optional[Union[str, Iterable[float]]]) -> Optional[Tuple[float, float, float, float]]:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
parts = [p.strip() for p in value.split(",") if p.strip()]
|
||||
if not parts:
|
||||
return None
|
||||
try:
|
||||
nums = [float(p) for p in parts]
|
||||
except ValueError:
|
||||
return None
|
||||
if len(nums) == 1:
|
||||
return (nums[0], nums[0], nums[0], nums[0])
|
||||
if len(nums) >= 4:
|
||||
return (nums[0], nums[1], nums[2], nums[3])
|
||||
return tuple(nums + [nums[-1]] * (4 - len(nums))) # pad to length 4
|
||||
try:
|
||||
vals = list(value)
|
||||
except TypeError:
|
||||
return None
|
||||
if not vals:
|
||||
return None
|
||||
vals = [float(v) for v in vals]
|
||||
if len(vals) == 1:
|
||||
return (vals[0], vals[0], vals[0], vals[0])
|
||||
if len(vals) >= 4:
|
||||
return (vals[0], vals[1], vals[2], vals[3])
|
||||
return tuple(vals + [vals[-1]] * (4 - len(vals)))
|
||||
|
||||
|
||||
def build_transform_chain(
|
||||
transform_specs: Iterable[object],
|
||||
*,
|
||||
backbone_name: str,
|
||||
augment: bool = True,
|
||||
mask_provider: Optional[Callable[[Image.Image, Optional[str]], np.ndarray]] = None,
|
||||
strict: bool = True,
|
||||
) -> transforms.Compose:
|
||||
"""
|
||||
Build an image transform pipeline from a list of transform specs plus the
|
||||
standard ToTensor + Normalize steps. This mirrors the V1 preprocessing
|
||||
but uses the explicit transform nodes from config.
|
||||
"""
|
||||
ops: list[Callable[[Image.Image], Image.Image]] = []
|
||||
for spec in transform_specs:
|
||||
transform_type = getattr(spec, "transform_type", None)
|
||||
params = getattr(spec, "params", None)
|
||||
if transform_type is None and isinstance(spec, dict):
|
||||
transform_type = spec.get("transformType") or spec.get("transform_type")
|
||||
params = spec
|
||||
params = params or {}
|
||||
|
||||
if transform_type == "resize":
|
||||
size = params.get("resizeSize", 256)
|
||||
ops.append(ResizeTransform(size=size))
|
||||
elif transform_type == "center_crop":
|
||||
size = params.get("centerCropSize", 224)
|
||||
ops.append(CenterCropTransform(size=size))
|
||||
elif transform_type == "jitter_bundle":
|
||||
if not augment:
|
||||
continue
|
||||
jitter = JitterBundleTransform(
|
||||
hflip=bool(params.get("jitterHFlip", True)),
|
||||
vflip=bool(params.get("jitterVFlip", True)),
|
||||
rotation_deg=int(params.get("jitterRotation", 15) or 0),
|
||||
color_jitter=_parse_color_jitter(params.get("jitterColor"))
|
||||
if params.get("jitterColorEnabled", True)
|
||||
else None,
|
||||
)
|
||||
ops.append(jitter)
|
||||
elif transform_type == "roi_crop":
|
||||
roi = ROICropTransform(
|
||||
mask_source=params.get("roiMaskSource", "gt"),
|
||||
mask_provider=mask_provider,
|
||||
scale=float(params.get("roiScale", 2.5)),
|
||||
target_size=(int(params.get("roiTargetSize", 224)), int(params.get("roiTargetSize", 224)))
|
||||
if params.get("roiTargetSize") is not None
|
||||
else None,
|
||||
fallback_to_original=bool(params.get("roiFallback", True)),
|
||||
)
|
||||
if roi.mask_provider is None and roi.mask_source == "unet":
|
||||
if strict:
|
||||
raise ValueError("ROI crop requires a mask provider for 'unet' source.")
|
||||
ops.append(roi)
|
||||
else:
|
||||
if strict:
|
||||
raise ValueError(f"Unsupported transform type: {transform_type!r}")
|
||||
|
||||
# Always end with tensor + normalize, using backbone defaults
|
||||
cfg = backbone_transform_config(backbone_name, augment=augment)
|
||||
ops.extend(
|
||||
[
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(mean=cfg.mean, std=cfg.std),
|
||||
]
|
||||
)
|
||||
return transforms.Compose(ops)
|
||||
Reference in New Issue
Block a user