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
|