File size: 7,150 Bytes
71d64bb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 | """
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
|