File size: 6,231 Bytes
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1eac9b8
8c867d9
 
 
 
 
 
 
 
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
"""Public Jev-style typed-decision corpora for training (DATA["text3"], train split only), decontaminated:
  * tasksource/procedural-typed-decisions (Apache-2.0): structured states, several typed questions, exact labels
  * n4ze3m/typed-decisions-synth (MIT): 149 domains, multi-question, soft teacher labels
  * tasksource/tasksource-jev-typed-decisions: real labels from 670 sources -- every source that overlaps a benchmark
    we report (AG News, emotion, Banking77, MASSIVE, XNLI, JevBench, typed-decisions) is excluded
Every candidate whose state matches an eval state (normalised prefix hash) is dropped as well.
"""
import hashlib, json, random, re
from datasets import load_dataset, get_dataset_config_names
from mmjev import Seg, options_of

rng = random.Random(7)
BANNED_SRC = re.compile(r"ag_?news|emotion|banking|massive|xnli|jev|typed|clinc|amazon_massive", re.I)


def norm(t):
    return hashlib.md5(re.sub(r"\s+", " ", str(t).lower())[:300].encode()).hexdigest()


EVAL_HASH = {norm(s.data) for v in DATA.values() for r in v if r["split"] == "eval"
             for s in r["state"] if s.kind == "text"}
out, dropped = [], {"contaminated": 0, "banned_source": 0, "bad": 0}


def convert_q(q):
    t, crit = q["type"], q.get("criteria")
    if t == "noul":
        c = crit or {}
        return {"type": "noul", "instructions": q["instructions"],
                "criteria": {"no": c.get("false", "no"), "yes": c.get("true", "yes")}}
    if t == "score":
        return {"type": "score", "instructions": q["instructions"], "criteria": [str(x) for x in crit]}
    if isinstance(crit, list):
        crit = {str(x): "" for x in crit}
    return {"type": "choice", "instructions": q["instructions"], "criteria": {str(k): str(v or "") for k, v in crit.items()}}


def soft_target(qq, ans):
    labels, _ = options_of(qq)
    if qq["type"] == "noul":
        p = float(ans.get("noul", ans.get("probabilities", {}).get("true", 0.5)) if isinstance(ans, dict) else ans)
        return [1 - p, p]
    pr = ans.get("probabilities", {}) if isinstance(ans, dict) else {}
    v = [float(pr.get(l, 0.0)) for l in labels]
    if sum(v) <= 0 and isinstance(ans, dict) and "choice" in ans and ans["choice"] in labels:
        v = [1.0 if l == ans["choice"] else 0.0 for l in labels]
    s = sum(v)
    return [x / s for x in v] if s > 0 else None


def add_case(task, state, qs, answers):
    if norm(state) in EVAL_HASH:
        dropped["contaminated"] += 1; return
    Q, Y, T = [], [], []
    for qid, q in qs.items():
        try:
            qq = convert_q(q); tg = soft_target(qq, answers[qid])
        except Exception:
            tg = None
        if tg is None or len(tg) < 2 or len(tg) > 80:
            continue
        Q.append(qq); T.append(tg); Y.append(int(max(range(len(tg)), key=lambda i: tg[i])))
    if not Q:
        dropped["bad"] += 1; return
    r = rec(task, "text", "train", [Seg("text", state[:4000])], Q[:6], Y[:6], T[:6])
    r["raw"] = None
    out.append(r)


def procedural(per_cfg=160):
    for cfg in get_dataset_config_names("tasksource/procedural-typed-decisions"):
        ds = load_dataset("tasksource/procedural-typed-decisions", cfg, split="train").shuffle(seed=7)
        for ex in ds.select(range(min(per_cfg, len(ds)))):
            add_case(f"procedural:{cfg}", ex["state"], json.loads(ex["questions"]), json.loads(ex["answers"]))


def synth(n=1800):
    ds = load_dataset("json", data_files="hf://datasets/n4ze3m/typed-decisions-synth/data/train.jsonl", split="train")
    for ex in ds.shuffle(seed=7).select(range(min(n, len(ds)))):
        st = ex["state"] if isinstance(ex["state"], str) else json.dumps(ex["state"], ensure_ascii=False)
        add_case("typed_synth", st, json.loads(ex["questions"]) if isinstance(ex["questions"], str) else ex["questions"],
                 json.loads(ex["teacher"]) if isinstance(ex["teacher"], str) else ex["teacher"])


def tasksource(n=2500):
    ds = load_dataset("tasksource/tasksource-jev-typed-decisions", split="train", streaming=True).shuffle(seed=7, buffer_size=20000)
    got = 0
    for ex in ds:
        if got >= n:
            break
        if BANNED_SRC.search(ex["source"] or ""):
            dropped["banned_source"] += 1; continue
        try:
            opts = json.loads(ex["options"]) if ex["options"] else []
            tgt = json.loads(ex["target"])
        except Exception:
            continue
        kind, st = ex["kind"], ex["state"] or ""
        if len(st) > 3000 or norm(st) in EVAL_HASH:
            dropped["contaminated" if norm(st) in EVAL_HASH else "bad"] += 1; continue
        if kind == "noul":
            p = float(tgt[-1] if isinstance(tgt, list) else tgt)
            q = {"type": "noul", "instructions": ex["question"]}; t = [1 - p, p]
        elif kind in ("choice", "score") and 2 <= len(opts) <= 40 and len(tgt) == len(opts):
            if kind == "score":
                q = {"type": "score", "instructions": ex["question"], "criteria": [str(o)[:120] for o in opts]}
            elif max(len(str(o)) for o in opts) > 60:   # long options: letter labels, text as criterion
                keys = [chr(65 + i) for i in range(len(opts))]
                q = {"type": "choice", "instructions": ex["question"], "criteria": {k: str(o)[:300] for k, o in zip(keys, opts)}}
            else:
                q = {"type": "choice", "instructions": ex["question"], "criteria": {str(o): "" for o in opts}}
                if len(q["criteria"]) != len(opts):
                    continue
            t = [float(x) for x in tgt]; s = sum(t)
            if s <= 0:
                continue
            t = [x / s for x in t]
        else:
            continue
        r = rec(f"tasksource:{ex['source'].split('/')[0]}", "text", "train", [Seg("text", st)], q,
                int(max(range(len(t)), key=lambda i: t[i])), [t])
        r["raw"] = None
        out.append(r); got += 1


for f in (procedural, synth):
    try:
        f()
        log(f"[text3] {f.__name__}: total {len(out)} dropped {dropped}")
    except Exception:
        import traceback
        log(f"[text3] {f.__name__} FAILED {traceback.format_exc()[-600:]}")
DATA["text3"] = out
log(f"[text3] DONE {len(out)} records, dropped {dropped}")