""" 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