Download main_method/code/augmentation.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 2.38 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/augmentation.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/augmentation.py
-
curl -L -o augmentation.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/augmentation.py
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 | |