Download code/build_text3.py from fnruha0921/omnijev-work: direct link, hf CLI and curl.
- Browser
- Download file 6.23 kB
-
https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/build_text3.py
- Command line
-
hf download hf://fnruha0921/omnijev-work/code/build_text3.py
-
curl -L -o build_text3.py https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/build_text3.py
6.23 kB
| """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}") | |