File size: 2,381 Bytes
71d64bb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import math,random
import torch
import config as C

def aug_dihedral(xc, yc):
    """One dihedral-8 transform per crop-batch, image and label together."""
    if torch.rand(()) > 0.5:
        xc, yc = xc.flip(-2), yc.flip(-2)
    if torch.rand(()) > 0.5:
        xc, yc = xc.flip(-1), yc.flip(-1)
    k = int(torch.randint(0, 4, ()))
    if k:
        xc, yc = torch.rot90(xc, k, (-2, -1)), torch.rot90(yc, k, (-2, -1))
    return xc.contiguous(), yc.contiguous()


def aug_temporal_subsample(xc, pos, p=0.5, drop_max=0.2, keep_min=24):
    """Drop up to `drop_max` of the frames, order preserved. Positions keep
    their original index values -> emulates missing acquisitions of an
    otherwise identical season. No-op with probability 1-p."""
    if torch.rand(()) > p:
        return xc, pos
    T = xc.shape[1]
    keep = int(round(T * (1.0 - random.uniform(0.0, drop_max))))
    keep = max(keep_min, min(T, keep))
    if keep >= T:
        return xc, pos
    idx = torch.randperm(T)[:keep].sort().values.to(xc.device)
    return (xc.index_select(1, idx).contiguous(),
            pos.index_select(1, idx).contiguous())


def aug_frame_dropout(xc, p=0.25, max_frames=2):
    """Zero 1-2 whole frames (0 == dataset mean after normalisation) —
    emulates cloudy/lost acquisitions."""
    if torch.rand(()) > p:
        return xc
    T = xc.shape[1]
    n = random.randint(1, max_frames)
    idx = torch.randperm(T)[:n].to(xc.device)
    xc = xc.clone()
    xc[:, idx] = 0.0
    return xc


def aug_spectral(xc, gain_std=0.05, bias_std=0.03, noise_std=0.01):
    """Per-sample, per-band gain/bias jitter + small gaussian noise (data is
    z-normalised, so these are fractions of one std)."""
    B, T, Ch, H, W = xc.shape
    dev = xc.device
    gain = 1.0 + gain_std * torch.randn(B, 1, Ch, 1, 1, device=dev)
    bias = bias_std * torch.randn(B, 1, Ch, 1, 1, device=dev)
    xc = xc * gain + bias
    if noise_std > 0:
        xc = xc + noise_std * torch.randn_like(xc)
    return xc.contiguous()


def lr_factor(step, total_steps, warmup_steps, min_ratio):
    """Linear warmup then cosine decay to min_ratio."""
    if step < warmup_steps:
        return (step + 1) / max(1, warmup_steps)
    p = (step - warmup_steps) / max(1, total_steps - warmup_steps)
    p = min(1.0, p)
    cos = 0.5 * (1.0 + math.cos(math.pi * p))
    return min_ratio + (1.0 - min_ratio) * cos