edge-predictor / scripts /augmentations.py
Yuvraj18's picture
Upload folder using huggingface_hub
c2dbb46 verified
Raw History Blame Contribute Delete
3.42 kB
"""Simple image augmentations for training."""
import numpy as np
import torch
def brightness_augment(
imgs: torch.Tensor,
coords: torch.Tensor,
masks: torch.Tensor,
*,
rng: np.random.Generator,
shift_range: float = 0.1,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Random additive brightness shift.
Parameters
----------
imgs : torch.Tensor
(W, *spatial) normalised images.
coords : torch.Tensor
(W, M, 3) node coordinates — passed through unchanged.
masks : torch.Tensor
(W, M) boolean masks — passed through unchanged.
rng : np.random.Generator
Random number generator.
shift_range : float
Maximum absolute shift sampled from [-shift_range, shift_range].
"""
shift = rng.uniform(-shift_range, shift_range)
return imgs + shift, coords, masks
def flip_augment(
imgs: torch.Tensor,
coords: torch.Tensor,
masks: torch.Tensor,
*,
rng: np.random.Generator,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Random spatial flip: samples uniformly from all 8 axis-aligned symmetries.
Each of Z, Y, X is independently flipped with probability 0.5,
giving 2^3 = 8 equally likely outcomes (including identity).
Creates a copy of coords so the underlying dataset is not mutated.
Only real (non-padded) coordinates are flipped; padding stays at zero.
"""
flip_mask = rng.random(3) < 0.5 # (Z, Y, X)
dims_to_flip = [1 + dim for dim, flip in enumerate(flip_mask) if flip]
if not dims_to_flip:
return imgs, coords, masks
imgs = imgs.flip(dims=dims_to_flip)
coords = coords.clone()
shape = imgs.shape[1:] # (Z, Y, X)
for dim in range(3):
if flip_mask[dim]:
coords[..., dim] = torch.where(
masks,
shape[dim] - coords[..., dim] - 1,
coords[..., dim],
)
return imgs, coords, masks
# === extra robustness augments =======================================
# Build invariance to the acquisition shifts a frozen model is weakest to
# (additive noise, blur, contrast/gamma) -> better held-out-group transfer.
# Intensity-only: coords/masks pass through.
import torch.nn.functional as _Fg
def noise_augment(imgs, coords, masks, *, rng, sigma_frac=0.5):
s = rng.uniform(0.0, sigma_frac) * float(imgs.std())
return imgs + torch.randn_like(imgs) * s, coords, masks
def contrast_augment(imgs, coords, masks, *, rng, lo=0.6, hi=1.6):
c = rng.uniform(lo, hi); mu = float(imgs.mean())
return (imgs - mu) * c + mu, coords, masks
def gamma_augment(imgs, coords, masks, *, rng, lo=0.6, hi=1.6):
g = rng.uniform(lo, hi)
return imgs.clamp(min=0).pow(g), coords, masks
def blur_augment(imgs, coords, masks, *, rng, max_k=3):
k = int(rng.integers(1, max_k + 1)) | 1
if k == 1:
return imgs, coords, masks
p = k // 2
# imgs: (W, Z, Y, X)
# Treat each temporal frame independently as a 3D volume.
x = imgs.float().unsqueeze(1) # (W, 1, Z, Y, X)
x = _Fg.pad(
x,
(p, p, p, p, p, p),
mode="replicate",
)
x = _Fg.avg_pool3d(
x,
kernel_size=k,
stride=1,
)
return x[:, 0], coords, masks
generalization_augments = [
brightness_augment, flip_augment,
noise_augment, contrast_augment, gamma_augment, blur_augment,
]