TSRDA / main_method /code /augmentation.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
2.38 kB
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