"""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, ]