Download main_method/code/dataset.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 7.15 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/dataset.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/dataset.py
-
curl -L -o dataset.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/dataset.py
7.15 kB
| """ | |
| PASTIS dataset loader for T-SRDA — OFFICIAL STCLN protocol. | |
| Each dataset item is ONE FULL 128x128 patch. Sub-cropping happens in the | |
| training loop, not here, exactly as the reference implementation does it: | |
| pretrain (pretraining_STCLN.py:123) for i in range(4): for j in range(4) | |
| finetune (finetuning_STCLN.py:185) for i,j in [[0,0],[1,1]] | |
| test (test_STCLN.py:52) for i in range(1) -> full patch | |
| Keeping the crop schedule in the loop is what makes the step counts line up | |
| with the paper: each (i, j) is a SEPARATE optimizer step, so pretrain does | |
| 124 batches x 16 crops = 1,984 steps/epoch, not 496 steps. | |
| Split builders: | |
| patches_in_folds([5]) -> 496 pretrain patch IDs | |
| train_patch_ids() -> 76 IDs (fold 1, duplicates kept) | |
| val_patch_ids() -> 76 IDs (fold 2, duplicates kept) | |
| test_patch_ids() -> 482 IDs (fold 4) | |
| """ | |
| import json | |
| from datetime import datetime | |
| import numpy as np | |
| import torch | |
| from torch.utils.data import Dataset | |
| import config as C | |
| REF_DATE = datetime(*map(int, C.REF_DATE_STR.split("-"))) | |
| # --------------------------------------------------------------------------- | |
| # Metadata + normalisation loaders (cached at module level) | |
| # --------------------------------------------------------------------------- | |
| def _date_to_days(date_int): | |
| s = str(int(date_int)) | |
| y, m, d = int(s[:4]), int(s[4:6]), int(s[6:8]) | |
| return (datetime(y, m, d) - REF_DATE).days | |
| def _load_meta(): | |
| raw = json.load(open(C.META_PATH)) | |
| info = {} | |
| for feat in raw["features"]: | |
| p = feat["properties"] | |
| pid = int(p.get("ID_PATCH", p.get("id_patch", -1))) | |
| fold = int(p.get("Fold", p.get("fold", 0))) | |
| dates_dict = p.get("dates-S2", {}) | |
| sorted_dates = sorted(dates_dict.items(), key=lambda x: int(x[0])) | |
| days = [_date_to_days(v) for _, v in sorted_dates] | |
| info[pid] = {"fold": fold, "dates": days} | |
| return info | |
| def _load_norm(): | |
| """Mean/std averaged across all folds; differs from fold-selected upstream normalization.""" | |
| raw = json.load(open(C.NORM_PATH)) | |
| if "mean" in raw: | |
| return (np.array(raw["mean"], dtype=np.float32), | |
| np.array(raw["std"], dtype=np.float32)) | |
| fold_keys = [k for k in raw if k.startswith("Fold_")] | |
| means = np.array([raw[k]["mean"] for k in fold_keys], dtype=np.float32) | |
| stds = np.array([raw[k]["std"] for k in fold_keys], dtype=np.float32) | |
| return means.mean(axis=0), stds.mean(axis=0) | |
| _META = None | |
| _NORM = None | |
| def get_meta(): | |
| global _META | |
| if _META is None: | |
| _META = _load_meta() | |
| return _META | |
| def get_norm(): | |
| global _NORM | |
| if _NORM is None: | |
| _NORM = _load_norm() | |
| return _NORM | |
| # --------------------------------------------------------------------------- | |
| # Split builders (official protocol) | |
| # --------------------------------------------------------------------------- | |
| def patches_in_folds(folds): | |
| """Sorted patch IDs belonging to the given folds.""" | |
| meta = get_meta() | |
| return sorted(p for p in meta if meta[p]["fold"] in folds) | |
| def pretrain_patch_ids(): | |
| """496 fold-5 patch IDs — the unlabelled pretraining pool.""" | |
| return patches_in_folds(C.PRETRAIN_FOLDS) | |
| def train_patch_ids(): | |
| """76 hardcoded fold-1 IDs. Duplicates are intentional and preserved.""" | |
| return list(C.TRAIN_PATCH_IDS) | |
| def val_patch_ids(): | |
| """76 hardcoded fold-2 IDs. Duplicates are intentional and preserved.""" | |
| return list(C.VAL_PATCH_IDS) | |
| def test_patch_ids(): | |
| """482 fold-4 patch IDs — evaluated as full 128x128 patches.""" | |
| return patches_in_folds(C.TEST_FOLDS) | |
| # --------------------------------------------------------------------------- | |
| # Patch-level Dataset | |
| # --------------------------------------------------------------------------- | |
| class PASTISPatchDataset(Dataset): | |
| """Returns ONE full 128x128 patch per __getitem__. | |
| __getitem__ -> (x, pos, days), y | |
| x (T, 10, 128, 128) float32, per-band normalised | |
| pos (T,) what the model is fed: arange(T) if USE_INDEX_POSITIONS | |
| else real day offsets | |
| days (T,) real day offsets from REF_DATE, always — for analysis | |
| y (128, 128) int64 semantic labels | |
| """ | |
| def __init__(self, patch_ids, load_target=True): | |
| self.ids = list(patch_ids) # duplicates preserved on purpose | |
| self.meta = get_meta() | |
| self.norm_mean, self.norm_std = get_norm() | |
| self.load_target = load_target | |
| def __len__(self): | |
| return len(self.ids) | |
| def __getitem__(self, idx): | |
| pid = self.ids[idx] | |
| x = np.load(C.DATA_S2_DIR / f"S2_{pid}.npy").astype(np.float32) | |
| x = (x - self.norm_mean[None, :, None, None]) \ | |
| / (self.norm_std[None, :, None, None] + 1e-8) | |
| T = x.shape[0] | |
| days = np.asarray(self.meta[pid]["dates"][:T], dtype=np.float32) | |
| if days.shape[0] < T: | |
| days = np.pad(days, (0, T - days.shape[0])) | |
| pos = np.arange(T, dtype=np.float32) if C.USE_INDEX_POSITIONS else days | |
| if self.load_target: | |
| target = np.load(C.ANNOT_DIR / f"TARGET_{pid}.npy") | |
| y = torch.from_numpy(target[0].astype(np.int64)) | |
| else: | |
| y = torch.zeros(C.PATCH_SIZE, C.PATCH_SIZE, dtype=torch.long) | |
| return (torch.from_numpy(x), | |
| torch.from_numpy(pos), | |
| torch.from_numpy(days)), y | |
| # --------------------------------------------------------------------------- | |
| # Collate (pads variable T to the batch max — official pad_collate behaviour) | |
| # --------------------------------------------------------------------------- | |
| def pad_collate(batch, pad_value=0.0): | |
| xs, ps, ds, ys = [], [], [], [] | |
| for (x, p, d), y in batch: | |
| xs.append(x); ps.append(p); ds.append(d); ys.append(y) | |
| max_T = max(x.shape[0] for x in xs) | |
| def pad(t, v=0.0): | |
| if t.shape[0] == max_T: | |
| return t | |
| z = torch.full((max_T - t.shape[0], *t.shape[1:]), v, dtype=t.dtype) | |
| return torch.cat([t, z], dim=0) | |
| return ((torch.stack([pad(x, pad_value) for x in xs]), | |
| torch.stack([pad(p) for p in ps]), | |
| torch.stack([pad(d) for d in ds])), | |
| torch.stack(ys)) | |
| # Backwards-compatible alias — the train scripts pass collate_fn=collate_fn. | |
| collate_fn = pad_collate | |
| # --------------------------------------------------------------------------- | |
| # Crop helper — the inner (i, j) loop of the official training scripts | |
| # --------------------------------------------------------------------------- | |
| def crop_ij(x, y, i, j, grid=None): | |
| """Slice crop (i, j) out of a batch of full patches. | |
| x: (B, T, C, H, W) y: (B, H, W) or None | |
| grid: 4 -> 32x32 crops (train/finetune); 1 -> the whole patch (test) | |
| Mirrors `split = input.shape[-1] // grid` in the reference scripts. | |
| """ | |
| grid = C.PRETRAIN_CROP_GRID if grid is None else grid | |
| s = x.shape[-1] // grid | |
| xc = x[:, :, :, i * s:(i + 1) * s, j * s:(j + 1) * s] | |
| yc = None if y is None else y[:, i * s:(i + 1) * s, j * s:(j + 1) * s] | |
| return xc, yc | |