judgetron / scripts /prepare_splits.py
Chucks90's picture
Fix Guardian adapter against real datasets; add remote runner
84d148b verified
Raw History Blame Contribute Delete
2.93 kB
"""Build leakage-safe manifests.
In-domain sources are split by task into train / cal / test_id.
The OOD source is never trained or calibrated on, except an optional small
task-disjoint slice (cal_ood) used only for the 'customer corrections' refit.
Example:
python scripts/prepare_splits.py \
--source rlbench=data/rlbench_fail/metadata_execution.jsonl \
--source bridge=data/bridge_fail/metadata_execution.jsonl \
--ood ur5=data/ur5_fail/metadata_execution.jsonl \
--ood-cal-frac 0.2 --out manifests/
"""
import argparse
import sys
from collections import Counter
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from judgecal.data import from_guardian_jsonl, grouped_split, save_manifest # noqa: E402
def parse(spec):
domain, path = spec.split("=", 1)
return domain, path
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--source", action="append", required=True, help="domain=path/to/metadata.jsonl")
ap.add_argument("--ood", action="append", required=True,
help="domain=path/to/metadata.jsonl (held-out domain); repeatable to pool "
"several upstream splits of the same domain into one OOD task pool")
ap.add_argument("--ood-cal-frac", type=float, default=0.2)
ap.add_argument("--out", default="manifests")
ap.add_argument("--seed", type=int, default=0)
a = ap.parse_args()
ind = []
for spec in a.source:
ind += from_guardian_jsonl(parse(spec)[1], domain=parse(spec)[0])
splits = grouped_split(ind, {"train": 0.7, "cal": 0.15, "test_id": 0.15}, seed=a.seed)
# Pooling several upstream splits of the OOD domain only widens its TASK pool, which is what the
# clustered bootstrap resamples. ur5 train alone carries 7 taskvars; after the cal_ood/test_ood
# split that leaves 5 task clusters in test_ood, and the headline OOD CIs become unreadable.
ood = []
for spec in a.ood:
d, path = parse(spec)
ood += from_guardian_jsonl(path, domain=d)
if a.ood_cal_frac > 0:
o = grouped_split(ood, {"cal_ood": a.ood_cal_frac, "test_ood": 1 - a.ood_cal_frac}, seed=a.seed)
splits.update(o)
else:
splits["test_ood"] = ood
for name, eps in splits.items():
save_manifest(eps, Path(a.out) / f"{name}.jsonl")
c = Counter(e.label for e in eps)
print(f"{name:9s} n={len(eps):6d} success={c[1]:6d} failure={c[0]:6d} "
f"tasks={len({e.task for e in eps}):4d} domains={sorted({e.domain for e in eps})}")
# Leakage check: no task in more than one split.
seen = {}
for name, eps in splits.items():
for t in {(e.domain, e.task) for e in eps}:
assert t not in seen, f"task {t} in both {seen[t]} and {name}"
seen[t] = name
print("leakage check passed: splits are task-disjoint")
if __name__ == "__main__":
main()