Download scripts/augmentations.py from Yuvraj18/edge-predictor: direct link, hf CLI and curl.
- Browser
- Download file 3.42 kB
-
https://huggingface.co/Yuvraj18/edge-predictor/resolve/main/scripts/augmentations.py
- Command line
-
hf download hf://Yuvraj18/edge-predictor/scripts/augmentations.py
-
curl -L -o augmentations.py https://huggingface.co/Yuvraj18/edge-predictor/resolve/main/scripts/augmentations.py
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, | |
| ] | |