"""Data audit for the decision SFT: dataset matrix, train/test overlap, tasksource provenance. python scripts/decisions/audit_data.py --publico DIR --mezcla DIR --out audit_data.json Overlap is measured at three levels between every training file (both stages) and every test file: pair identical (state, question) after canonical JSON state identical state fp identical structural fingerprint of the state (non-text leaves; survives translation) For tasksource it also recovers each question's full source, group and variant from the original parquet (by id) and checks source / group / normalised-text overlap between train and test. """ from __future__ import annotations import argparse import collections import hashlib import json import re import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parent)) from prepare_mix import fingerprint # noqa: E402 TASKSOURCE = "tasksource/tasksource-jev-typed-decisions" def jsonl(p): return [json.loads(x) for x in open(p, encoding="utf-8") if x.strip()] def canon(x): return x if isinstance(x, str) else json.dumps(x, ensure_ascii=False, sort_keys=True) def h(s): return hashlib.sha1(s.encode()).hexdigest() def norm_text(s): return re.sub(r"\W+", " ", str(s).lower()).strip() def keys(cases): pair, state, fp, text = set(), set(), set(), set() for c in cases: st = canon(c["state"]) state.add(h(st)) text.add(h(norm_text(st))) if isinstance(c["state"], dict): fp.add(fingerprint(c["state"])) for q in c["questions"].values(): pair.add(h(st + "\x00" + canon(q))) return {"pair": pair, "state": state, "fp": fp, "text": text} def describe(cases): qs = [q for c in cases for q in c["questions"].values()] return {"cases": len(cases), "questions": len(qs), "types": dict(collections.Counter(q["type"] for q in qs)), "groups": len({c.get("grupo") for c in cases}), "langs": dict(collections.Counter(c.get("lang") for c in cases))} def tasksource_meta(ids): """id -> (source, group_id, variant, license) from the original parquet files.""" import pyarrow.parquet as pq from huggingface_hub import HfApi, hf_hub_download want = {i[3:] for i in ids} # strip "ts-" out = {} for f in sorted(s.rfilename for s in HfApi().dataset_info(TASKSOURCE).siblings if s.rfilename.endswith(".parquet")): pf = pq.ParquetFile(hf_hub_download(TASKSOURCE, f, repo_type="dataset")) for b in pf.iter_batches(columns=["id", "source", "group_id", "variant", "license", "question", "state"], batch_size=100000): for r in b.to_pylist(): if r["id"] in want: out["ts-" + r["id"]] = r return out def main(): ap = argparse.ArgumentParser() ap.add_argument("--publico", type=Path, required=True) ap.add_argument("--mezcla", type=Path, required=True) ap.add_argument("--out", type=Path, required=True) a = ap.parse_args() P, M = a.publico, a.mezcla sets = { # stage 1 (v2) "s1:typed_train_en": P / "typed_train_en.jsonl", "s1:typed_train_es": P / "typed_train_es.jsonl", "s1:jev_train": P / "jev_train.jsonl", "val:typed_val_en": P / "typed_val_en.jsonl", "val:typed_val_es": P / "typed_val_es.jsonl", "val:jev_val": P / "jev_val.jsonl", # stage 2 (v3) **{"s2:" + p.stem: p for p in sorted(M.glob("train_*.jsonl"))}, **{"val2:" + p.stem: p for p in sorted(M.glob("val_*.jsonl"))}, # tests "test:typed_en": P / "typed_test_en.jsonl", "test:typed_es": P / "typed_test_es.jsonl", "test:telepatia_es": M / "test_telepatia_es.jsonl", "test:tasksource": M / "test_tasksource.jsonl", } data = {k: jsonl(v) for k, v in sets.items()} res = {"files": {k: dict(describe(v), path=str(sets[k].name), sha256=hashlib.sha256(sets[k].read_bytes()).hexdigest()) for k, v in data.items()}} K = {k: keys(v) for k, v in data.items()} overlap = {} for t in [k for k in data if k.startswith("test:")]: for s in [k for k in data if not k.startswith("test:")]: row = {lvl: len(K[t][lvl] & K[s][lvl]) for lvl in ("pair", "state", "fp", "text")} if any(row.values()): overlap[t + " x " + s] = row res["overlap_test_vs_train_val"] = overlap # tasksource provenance ts_ids = [c["id"] for k in ("s2:train_tasksource", "val2:val_tasksource", "test:tasksource") for c in data[k]] meta = tasksource_meta(ts_ids) def fam(k): rows = [meta.get(c["id"]) for c in data[k]] return rows tr, te, va = fam("s2:train_tasksource"), fam("test:tasksource"), fam("val2:val_tasksource") missing = sum(r is None for r in tr + te + va) src_tr = {r["source"] for r in tr if r} grp_tr = {r["group_id"] for r in tr if r} q_tr = {h(norm_text(r["question"])) for r in tr if r} st_tr = {h(norm_text(r["state"])) for r in tr if r} per_item = [] for c, r in zip(data["test:tasksource"], te): per_item.append({"id": c["id"], "source": r["source"], "family": r["source"].split("/")[0], "variant": r["variant"], "license": r["license"], "source_in_train": r["source"] in src_tr, "family_in_train": r["source"].split("/")[0] in {s.split("/")[0] for s in src_tr}, "group_in_train": r["group_id"] in grp_tr, "question_text_in_train": h(norm_text(r["question"])) in q_tr, "state_text_in_train": h(norm_text(r["state"])) in st_tr}) agg = {k: sum(x[k] for x in per_item) for k in ("source_in_train", "family_in_train", "group_in_train", "question_text_in_train", "state_text_in_train")} res["tasksource"] = {"meta_missing": missing, "train_sources": len(src_tr), "test_sources": len({x["source"] for x in per_item}), "test_families": len({x["family"] for x in per_item}), "test_items": len(per_item), "test_items_with": agg, "test_sources_not_in_train": sorted({x["source"] for x in per_item if not x["source_in_train"]}), "test_variants": dict(collections.Counter(x["variant"] for x in per_item))} (a.out.parent / "tasksource_test_items.jsonl").write_text("".join(json.dumps(x) + "\n" for x in per_item)) # typed workflows in every set res["typed_workflows"] = {k: sorted({c.get("grupo") for c in v}) for k, v in data.items() if "typed" in k or "telepatia" in k} a.out.write_text(json.dumps(res, indent=1, ensure_ascii=False) + "\n") print(json.dumps({k: res[k] for k in ("overlap_test_vs_train_val", "tasksource")}, indent=1)[:6000]) if __name__ == "__main__": main()