Ines-1 / training /code /audit_data.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
6.95 kB
"""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()