File size: 11,801 Bytes
b792420
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
"""Stage-3 public corpora (all HF, train splits only) -> MM-Jev records. Same record schema / caching as build_text4.

  * ZefanCai/Open-Jev snake-v1 (CC0): train (6x6) + ood (8x8) groups for training, test groups held out as snake_v1_* eval
  * OpenAGILab/Jev-dataset (CC0): the *-control-v1 families + drone / browser control (sources not already in Open-Jev v2)
  * AlexWortega/openjev-data: distilled hard decisions, tool calls, long-doc, instruction-following (3-way NLI -> noul)
  * moganai/mogan-decision-distill: MCQ with teacher distributions (GSM8K, MedMCQA, MMLU-aux, OpenBookQA, SciQ, SuperGPQA)
  * helmo/synthetic-typed-decisions, vagmi/jevlite_dataset, chand1012/laya-task-routing-synthetic-v2,
    pngwn/typed-decisions-v2-system-one: typed-decision corpora
  * multilingual typed decisions: canbingol/mmlu_typed_decision (tr), ahmetege/turkish_jev_noul (tr),
    fukayatti0/jev-japanese-judgment (ja), servronix/laya-thai-trainset (th), ai-simonsk13/id-typed-decisions (id)

Benchmark sources we report are dropped by name (Banking77 / CLINC / MASSIVE / XNLI / AG News / emotion / typed-decisions).
"""
import gzip, json, random, re
from datasets import load_dataset
from huggingface_hub import hf_hub_download
from build_text4 import _rec, _norm, _parse, row_question, grouped_rows, MAX_OPTS
from mmjev import Seg

BANNED = re.compile(r"banking|clinc|massive|xnli|ag_?news|emotion|typed_decisions", re.I)
SNAKE_Q = "Choose a direction to keep the snake alive and collect food."


# ------------------------------------------------------------------ typed-decision schema (state / questions / gold)
def typed_question(q, g):
    """{"type","instructions","criteria"} + gold {"label"/"probabilities"/"noul"} -> (question, target) or None."""
    t, crit = q.get("type"), q.get("criteria")
    g = g or {}
    if t == "noul":
        p = g.get("noul")
        if p is None:
            pr = g.get("probabilities") or {}
            p = pr.get("true", pr.get("yes"))
        if p is None and "label" in g:
            p = 1.0 if str(g["label"]).lower() in ("true", "yes", "1") else 0.0
        if p is None:
            return None
        qq = {"type": "noul", "instructions": q["instructions"]}
        if isinstance(crit, dict) and ({"true", "false"} <= set(crit) or {"yes", "no"} <= set(crit)):
            qq["criteria"] = {"no": crit.get("false", crit.get("no", "no")), "yes": crit.get("true", crit.get("yes", "yes"))}
        return qq, [1 - float(p), float(p)]
    pr = g.get("probabilities") or {}
    if t == "score":
        levels = list(crit) if isinstance(crit, list) else list((crit or {}).values()) if isinstance(crit, dict) else []
        keys = list(crit) if isinstance(crit, dict) else [str(i) for i in range(len(levels))]
        if not (2 <= len(levels) <= MAX_OPTS):
            return None
        v = [float(pr.get(k, pr.get(str(i), 0.0))) for i, k in enumerate(keys)]
        if sum(v) <= 0 and "label" in g:
            lab = str(g["label"])
            v = [1.0 if (lab == k or lab == str(i)) else 0.0 for i, k in enumerate(keys)]
        v = _norm(v)
        return ({"type": "score", "instructions": q["instructions"], "criteria": [str(x)[:160] for x in levels]}, v) if v else None
    crit = crit if isinstance(crit, dict) else {str(c): "" for c in (crit or [])}
    keys = list(crit)
    if not (2 <= len(keys) <= MAX_OPTS):
        return None
    v = [float(pr.get(k, 0.0)) for k in keys]
    if sum(v) <= 0 and str(g.get("label")) in keys:
        v = [1.0 if k == str(g["label"]) else 0.0 for k in keys]
    v = _norm(v)
    return ({"type": "choice", "instructions": q["instructions"], "criteria": {k: str(x or "")[:300] for k, x in crit.items()}}, v) if v else None


def typed_rows(rows, task, n, rng, state="state", questions="questions", gold="gold", task_of=None):
    rows = list(rows); rng.shuffle(rows)
    out = []
    for ex in rows:
        if len(out) >= n:
            break
        try:
            qs, gs = _parse(ex[questions]), _parse(ex[gold])
        except Exception:
            continue
        Q, T = [], []
        for qid, q in list(qs.items())[:6]:
            got = typed_question(q, gs.get(qid))
            if got:
                Q.append(got[0]); T.append(got[1])
        if Q:
            st = ex[state] if isinstance(ex[state], str) else json.dumps(ex[state], ensure_ascii=False)
            out.append(_rec(task_of(ex) if task_of else task, st, Q, T))
    return out


# ------------------------------------------------------------------ builders
def openjev_snake(seed=21):
    """Every snake-v1 group of the train + ood splits (4 questions per board); test boards -> held-out eval."""
    out = []
    for split in ("train", "ood", "test"):
        ds = load_dataset("ZefanCai/Open-Jev", "release-v2-redistributable", split=split)
        rows = [r for r in ds if r["source"] == "snake-v1"]
        if split != "test":
            out += grouped_rows(rows, "snake_v1", 10 ** 9, random.Random(seed))
            continue
        for r in rows:                                   # eval: one record per question family
            got = row_question(r["kind"], r["question"], _parse(r["options"]), _parse(r["target"]))
            if not got:
                continue
            e = _rec("snake_v1_action" if r["kind"] == "choice" else "snake_v1_collision", r["state_json"], [got[0]], [got[1]])
            e["split"] = "eval"
            out.append(e)
    return out


def openagi_controls(n_per=500, n_expansion=2500, seed=22):
    fams = ["amount-extraction", "citation", "context-retention", "email-selection", "entity-alignment", "ir", "mailroom",
            "phone-extraction", "silent-failure", "sponsor-segment"]
    rng, out = random.Random(seed), []
    for f in fams:
        ds = load_dataset("parquet", data_files=f"hf://datasets/OpenAGILab/Jev-dataset/data/{f}-control-v1/train-*.parquet", split="train")
        out += grouped_rows(ds.to_list(), f"openagi:{f}", n_per, rng)
    ds = load_dataset("parquet", data_files="hf://datasets/OpenAGILab/Jev-dataset/data/browser-drone-expansion-v1-redistributable/train-*.parquet",
                      split="train")
    rows = [r for r in ds.to_list() if r["source"] in ("drone-control-v1", "browser-control-v1")]
    out += grouped_rows(rows, "openagi:drone_browser", n_expansion, rng)
    return out


AWD_FILES = {"distill.jsonl.gz": 2500, "distill_hard.jsonl.gz": 2500, "hardfmt.jsonl.gz": 1000, "agentic.jsonl.gz": 1200,
             "longdoc.jsonl.gz": 800, "ifcomplex.jsonl.gz": 1200}


def openjev_data(seed=23, scan=60000):
    """3-way NLI rows (0 contradiction / 1 entailment / 2 neutral) -> noul 'is the claim supported by the state'."""
    rng, out = random.Random(seed), []
    for f, n in AWD_FILES.items():
        rows = []
        with gzip.open(hf_hub_download("AlexWortega/openjev-data", f, repo_type="dataset"), "rt") as fh:
            for i, line in enumerate(fh):
                if i >= scan:
                    break
                r = json.loads(line)
                if r.get("image") or BANNED.search(r.get("source") or "") or len(r["premise"]) > 3500:
                    continue
                rows.append(r)
        for r in rng.sample(rows, min(n, len(rows))):
            p = 1.0 if r["label"] == 1 else 0.0
            q = {"type": "noul", "instructions": f"Given the state, is this claim correct? Claim: {r['hypothesis'][:600]}"}
            out.append(_rec(f"openjev_data:{r['source']}", r["premise"], [q], [[1 - p, p]]))
    return out


def mogan(n=5000, seed=24):
    files = ["gsm8k", "medmcqa", "mmlu_auxiliary_train", "openbookqa", "sciq", "supergpqa"]
    rng, out = random.Random(seed), []
    for f in files:
        ds = load_dataset("json", data_files=f"hf://datasets/moganai/mogan-decision-distill/data/{f}.jsonl.gz", split="train")
        for ex in ds.shuffle(seed=seed).select(range(min(len(ds), n // len(files)))):
            q, tp = ex["question"], ex["teacher_probs"] or {}
            keys = list(q["criteria"])
            v = _norm([float(tp.get(k) or 0.0) for k in keys])
            if not v or not (2 <= len(keys) <= MAX_OPTS):
                continue
            st = ex["state"] if isinstance(ex["state"], str) else json.dumps(ex["state"], ensure_ascii=False)
            out.append(_rec(f"mogan:{f}", st, [{"type": "choice", "instructions": q["instructions"],
                                                  "criteria": {k: str(c)[:300] for k, c in q["criteria"].items()}}], [v]))
    return out


def typed_corpora(seed=25):
    rng, out = random.Random(seed), []
    ds = load_dataset("json", data_files="hf://datasets/helmo/synthetic-typed-decisions/synthetic_train.jsonl", split="train")
    out += typed_rows(ds, "helmo_typed", 2500, rng)
    ds = load_dataset("json", data_files="hf://datasets/chand1012/laya-task-routing-synthetic-v2/unified/train.jsonl", split="train")
    out += typed_rows(ds, "task_routing", 1500, rng)
    ds = load_dataset("json", data_files="hf://datasets/vagmi/jevlite_dataset/synth.train.jsonl", split="train")
    for ex in ds.shuffle(seed=seed).select(range(min(len(ds), 2000))):
        got = row_question(ex["type"], ex["question"], ex["options"], ex["label"])
        if got:
            out.append(_rec("jevlite", ex["state"], [got[0]], [got[1]]))
    ds = load_dataset("pngwn/typed-decisions-v2-system-one", split="train").shuffle(seed=seed)
    for ex in ds.select(range(min(len(ds), 1500))):
        k = len(ex["options"])
        kind = "score" if ex["ordered"] else ("noul" if ex["question_type"] == "noul" else "choice")
        tgt = [1.0 if i == ex["answer_index"] else 0.0 for i in range(k)]
        if kind == "noul":
            opts = [str(o).lower() for o in ex["options"]]
            yes = opts.index("yes") if "yes" in opts else 1
            kind, tgt = "noul", [1.0 - tgt[yes], tgt[yes]]
            got = ({"type": "noul", "instructions": ex["question"]}, tgt)
        else:
            got = row_question(kind, ex["question"], ex["options"], tgt)
        if got:
            out.append(_rec(f"pngwn:{ex['task']}", ex["state"], [got[0]], [got[1]]))
    return out


def multilingual(seed=26):
    rng, out = random.Random(seed), []
    out += typed_rows(load_dataset("canbingol/mmlu_typed_decision", split="train"), "tr_mmlu", 1200, rng)
    out += typed_rows(load_dataset("ahmetege/turkish_jev_noul", split="train"), "tr_noul", 1000, rng)
    th = load_dataset("json", data_files="hf://datasets/servronix/laya-thai-trainset/trainset.jsonl", split="train")
    th = th.filter(lambda e: not BANNED.search(e["workflow"] or ""))
    out += typed_rows(th, "th", 1500, rng, task_of=lambda e: f"th:{e['workflow']}")
    out += typed_rows(load_dataset("json", data_files="hf://datasets/ai-simonsk13/id-typed-decisions/train.jsonl", split="train"),
                      "id_typed", 10 ** 9, rng)
    ja = load_dataset("fukayatti0/jev-japanese-judgment", split="train").shuffle(seed=seed)
    for ex in ja.select(range(min(len(ja), 1500))):
        c = [str(x) for x in ex["candidates"]]
        if not (2 <= len(c) <= MAX_OPTS) or len(set(c)) != len(c) or not (0 <= ex["label"] < len(c)):
            continue
        st = (ex["context"] + "\n" if ex["context"] else "") + ex["question"]
        out.append(_rec(f"ja:{ex['source_dataset']}", st, [{"type": "choice", "instructions": "ζœ€γ‚‚ι©εˆ‡γͺη­”γˆγ‚’ιΈγ‚“γ§γγ γ•γ„γ€‚",
                                                            "criteria": {x: "" for x in c}}],
                        [[1.0 if i == ex["label"] else 0.0 for i in range(len(c))]]))
    return out


BUILDERS = {"openjev_snake": openjev_snake, "openagi_controls": openagi_controls, "openjev_data": openjev_data,
            "mogan": mogan, "typed_corpora": typed_corpora, "multilingual": multilingual}