File size: 3,421 Bytes
fdae75f
 
 
 
 
 
 
 
 
 
 
 
 
 
c2dbb46
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fdae75f
 
 
 
 
 
 
 
 
 
 
c2dbb46
fdae75f
c2dbb46
 
 
 
 
 
 
fdae75f
 
 
 
 
 
 
 
c2dbb46
fdae75f
c2dbb46
 
 
 
 
 
 
fdae75f
c2dbb46
fdae75f
 
c2dbb46
 
 
fdae75f
 
 
 
 
 
 
c2dbb46
fdae75f
 
 
 
 
 
 
 
 
 
 
 
 
 
c2dbb46
 
 
 
 
 
 
 
 
 
fdae75f
c2dbb46
 
fdae75f
 
 
c2dbb46
fdae75f
 
c2dbb46
 
fdae75f
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
"""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,
]