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