File size: 6,948 Bytes
61b6fb9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
"""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()