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