TSRDA / check_data.py
Dhruv1000's picture
Upload TSRDA project: code, checkpoints, logs, results
a062862 verified
Raw History Blame Contribute Delete
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()