TSRDA / main_method /code /dataset.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
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