Download check_data.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 6.53 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/check_data.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/check_data.py
-
curl -L -o check_data.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/check_data.py
6.53 kB
| """Sanity check: T-SRDA data pipeline against the OFFICIAL STCLN protocol. | |
| Asserts every number in the protocol table, so a silent path/split regression | |
| fails here instead of 40 GPU-hours later. | |
| """ | |
| import sys, os | |
| sys.path.insert(0, os.path.dirname(__file__)) | |
| import numpy as np | |
| from collections import Counter | |
| from torch.utils.data import DataLoader | |
| import config as C | |
| from dataset import (PASTISPatchDataset, get_meta, pad_collate, crop_ij, | |
| patches_in_folds, pretrain_patch_ids, | |
| train_patch_ids, val_patch_ids, test_patch_ids) | |
| def check(label, got, want): | |
| ok = got == want | |
| print(f" {'OK ' if ok else 'FAIL'} {label:<44} {got}" | |
| f"{'' if ok else f' (expected {want})'}") | |
| assert ok, f"{label}: got {got}, expected {want}" | |
| def main(): | |
| print("=" * 66) | |
| print(" T-SRDA / PASTIS protocol check (official STCLN)") | |
| print("=" * 66) | |
| print(f" PASTIS_ROOT: {C.PASTIS_ROOT}\n") | |
| meta = get_meta() | |
| folds = Counter(v["fold"] for v in meta.values()) | |
| print("Fold inventory") | |
| check("total patches in metadata", len(meta), 2433) | |
| for f in sorted(folds): | |
| role = ("pretrain" if f in C.PRETRAIN_FOLDS else | |
| "test" if f in C.TEST_FOLDS else "unused/labelled") | |
| print(f" fold {f}: {folds[f]:>4} patches [{role}]") | |
| print("\nPretrain split (fold 5, unlabelled)") | |
| pre = pretrain_patch_ids() | |
| check("pretrain patches", len(pre), 496) | |
| check("pretrain crop-steps/epoch", | |
| (len(pre) // C.PRE_BATCH) * C.CROPS_PER_PATCH, 1984) | |
| check("pretrain batches/epoch", len(pre) // C.PRE_BATCH, 124) | |
| print("\nFinetune splits (hardcoded patch IDs)") | |
| tr, va = train_patch_ids(), val_patch_ids() | |
| check("train patch IDs", len(tr), 76) | |
| check("train unique IDs", len(set(tr)), 72) | |
| check("val patch IDs", len(va), 76) | |
| check("val unique IDs", len(set(va)), 71) | |
| check("train crops (76 x FT_CROP_IJ)", len(tr) * len(C.FT_CROP_IJ), 152) | |
| check("val crops", len(va) * len(C.FT_CROP_IJ), 152) | |
| check("finetune steps/epoch", | |
| -(-len(tr) // C.FT_BATCH) * len(C.FT_CROP_IJ), 76) | |
| check("train IDs all in fold 1", | |
| sorted({meta[p]["fold"] for p in tr}), [1]) | |
| check("val IDs all in fold 2", | |
| sorted({meta[p]["fold"] for p in va}), [2]) | |
| print(f" train duplicates: " | |
| f"{sorted(k for k, v in Counter(tr).items() if v > 1)}") | |
| print(f" val duplicates: " | |
| f"{sorted(k for k, v in Counter(va).items() if v > 1)}") | |
| print("\nTest split (fold 4, full patches)") | |
| te = test_patch_ids() | |
| check("test patches", len(te), 482) | |
| check("test crop-equivalent", len(te) * C.CROPS_PER_PATCH, 7712) | |
| check("full-patch evaluation", C.TEST_FULL_PATCH, True) | |
| for b in (1, 2, 4): | |
| print(f" eval batches at batch {b}: {-(-len(te) // b)}") | |
| print("\nDisjointness") | |
| S = {"pretrain": set(pre), "train": set(tr), "val": set(va), "test": set(te)} | |
| for a in S: | |
| for b in S: | |
| if a < b: | |
| check(f"{a} ∩ {b}", len(S[a] & S[b]), 0) | |
| unused = set(meta) - set().union(*S.values()) | |
| by_fold = Counter(meta[p]["fold"] for p in unused) | |
| check("fold 3 entirely unused", by_fold[3], folds[3]) | |
| check("no fold-4 or fold-5 patch left out", by_fold[4] + by_fold[5], 0) | |
| print(f" unused patches: {len(unused)} of {len(meta)}") | |
| for f in sorted(by_fold): | |
| print(f" fold {f}: {by_fold[f]:>4} unused of {folds[f]:>4}" | |
| f"{' (only 72/71 labelled IDs are referenced)' if f in (1, 2) else ''}") | |
| print("\nClass nomenclature") | |
| OFFICIAL_CLASSES = [ | |
| "Background", "Meadow", "Soft winter wheat", "Corn", "Winter barley", | |
| "Winter rapeseed", "Spring barley", "Sunflower", "Grapevine", "Beet", | |
| "Winter triticale", "Winter durum wheat", "Fruits, vegetables, flowers", | |
| "Potatoes", "Leguminous fodder", "Soybeans", "Orchard", "Mixed cereal", | |
| "Sorghum", "Void label", | |
| ] | |
| check("class names match the official PASTIS list", | |
| [C.PASTIS_CLASSES[i] for i in range(20)], OFFICIAL_CLASSES) | |
| ref = C.REPO_ROOT.parent / "test" / "PhenoProto-SSL" | |
| if (ref / "splits.py").is_file(): | |
| sys.path.insert(0, str(ref)) | |
| try: | |
| import splits as _pp | |
| check("matches PhenoProto-SSL splits.PASTIS_CLASSES", | |
| _pp.PASTIS_CLASSES, C.PASTIS_CLASSES) | |
| except Exception as e: | |
| print(f" (cross-check against PhenoProto skipped: {e})") | |
| finally: | |
| sys.path.pop(0) | |
| else: | |
| print(" (PhenoProto-SSL not present — cross-check skipped)") | |
| print("\nSample patch") | |
| ds = PASTISPatchDataset(tr[:1]) | |
| (x, pos, days), y = ds[0] | |
| print(f" pid {tr[0]}") | |
| print(f" x {tuple(x.shape)} {x.dtype} " | |
| f"range=[{x.min():.2f}, {x.max():.2f}]") | |
| print(f" y {tuple(y.shape)} classes=" | |
| f"{sorted(y.unique().tolist())}") | |
| print(f" pos (model in) {pos[:4].tolist()} ... " | |
| f"{'index' if C.USE_INDEX_POSITIONS else 'day offsets'}") | |
| print(f" days (real) {days[:4].tolist()} ...") | |
| check("x is a full patch", tuple(x.shape[-2:]), (128, 128)) | |
| if C.USE_INDEX_POSITIONS: | |
| check("pos == arange(T)", pos.tolist(), list(range(x.shape[0]))) | |
| print("\nBatching + crop loop") | |
| dl = DataLoader(PASTISPatchDataset(tr), batch_size=C.FT_BATCH, | |
| shuffle=False, num_workers=0, collate_fn=pad_collate) | |
| (xb, pb, db), yb = next(iter(dl)) | |
| print(f" batch x {tuple(xb.shape)}") | |
| print(f" batch y {tuple(yb.shape)}") | |
| xc, yc = crop_ij(xb, yb, 1, 1, C.PRETRAIN_CROP_GRID) | |
| print(f" crop (1,1) x {tuple(xc.shape)}") | |
| print(f" crop (1,1) y {tuple(yc.shape)}") | |
| check("finetune crop size", tuple(xc.shape[-2:]), (32, 32)) | |
| xf, yf = crop_ij(xb, yb, 0, 0, 1) | |
| check("test window (grid=1)", tuple(xf.shape[-2:]), (128, 128)) | |
| T = np.array([len(meta[p]["dates"]) for p in pre]) | |
| print(f"\nTemporal length (fold 5): min={T.min()} max={T.max()} " | |
| f"mean={T.mean():.1f}") | |
| T4 = np.array([len(meta[p]["dates"]) for p in te]) | |
| print(f"Temporal length (fold 4): min={T4.min()} max={T4.max()} " | |
| f"mean={T4.mean():.1f}") | |
| print("\n" + "=" * 66) | |
| print(" All protocol assertions passed.") | |
| print("=" * 66) | |
| if __name__ == "__main__": | |
| main() | |