Download check_split.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 3.19 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/check_split.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/check_split.py
-
curl -L -o check_split.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/check_split.py
3.19 kB
| """Diagnostic: class distribution in the official 152 train / 152 val crops.""" | |
| import sys, os | |
| sys.path.insert(0, os.path.dirname(__file__)) | |
| import numpy as np | |
| from collections import Counter | |
| import config as C | |
| from dataset import train_patch_ids, val_patch_ids, get_meta | |
| PASTIS_CLASSES = C.PASTIS_CLASSES # single source of truth — see config.py | |
| def crop_targets(patch_ids): | |
| """Yield the (i, j) label crops the finetune loop actually trains on.""" | |
| s = C.PATCH_SIZE // C.PRETRAIN_CROP_GRID | |
| for pid in patch_ids: # duplicates counted twice | |
| tgt = np.load(C.ANNOT_DIR / f"TARGET_{pid}.npy")[0] | |
| for (i, j) in C.FT_CROP_IJ: | |
| yield tgt[i * s:(i + 1) * s, j * s:(j + 1) * s] | |
| def main(): | |
| meta = get_meta() | |
| tr, va = train_patch_ids(), val_patch_ids() | |
| print(f"Train: {len(tr)} patches x {len(C.FT_CROP_IJ)} crops = " | |
| f"{len(tr) * len(C.FT_CROP_IJ)} crops (fold " | |
| f"{sorted({meta[p]['fold'] for p in tr})})") | |
| print(f"Val: {len(va)} patches x {len(C.FT_CROP_IJ)} crops = " | |
| f"{len(va) * len(C.FT_CROP_IJ)} crops (fold " | |
| f"{sorted({meta[p]['fold'] for p in va})})") | |
| print(f"Crop positions: {C.FT_CROP_IJ}\n") | |
| stats = {} | |
| for split_name, ids in [("TRAIN", tr), ("VAL", va)]: | |
| px = Counter() | |
| ncrop = Counter() | |
| total = 0 | |
| for sem in crop_targets(ids): | |
| total += sem.size | |
| for c in np.unique(sem): | |
| if 0 < c < 19: | |
| ncrop[int(c)] += 1 | |
| for c in range(20): | |
| px[c] += int((sem == c).sum()) | |
| stats[split_name] = (px, ncrop, total) | |
| print(f"{'='*66}") | |
| print(f" {split_name} — {len(ids) * len(C.FT_CROP_IJ)} crops, " | |
| f"{total:,} pixels") | |
| print(f"{'='*66}") | |
| print(f"{'Cls':>4} {'Name':<28} {'Pixels':>10} {'%':>7} {'Crops':>6} Status") | |
| print("-" * 66) | |
| for c in range(20): | |
| n = px.get(c, 0) | |
| pct = 100.0 * n / total if total else 0.0 | |
| if c in (0, 19): | |
| status = "ignore" | |
| elif n == 0: | |
| status = "!! DEAD" | |
| elif n < 500: | |
| status = "! RARE" | |
| elif n < 2000: | |
| status = "~ sparse" | |
| else: | |
| status = "ok" | |
| print(f"{c:>4} {PASTIS_CLASSES[c]:<28} {n:>10,} {pct:>6.2f}% " | |
| f"{ncrop.get(c, 0):>6} {status}") | |
| print() | |
| # train/val agreement on the scored classes | |
| trp, _, trt = stats["TRAIN"] | |
| vap, _, vat = stats["VAL"] | |
| print("=" * 66) | |
| print(" Train vs val share per scored class (ratio of pixel %)") | |
| print("=" * 66) | |
| print(f"{'Cls':>4} {'Name':<28} {'train %':>9} {'val %':>9} {'ratio':>8}") | |
| print("-" * 66) | |
| for c in range(1, 19): | |
| a = 100.0 * trp[c] / trt if trt else 0.0 | |
| b = 100.0 * vap[c] / vat if vat else 0.0 | |
| r = (a / b) if b > 0 else float("inf") | |
| flag = " <-- skewed" if (r > 2 or r < 0.5) else "" | |
| print(f"{c:>4} {PASTIS_CLASSES[c]:<28} {a:>8.2f}% {b:>8.2f}% " | |
| f"{r:>8.2f}{flag}") | |
| if __name__ == "__main__": | |
| main() | |