omnijev-work / code /build_text3.py
fnruha0921's picture
stage-2 multimodal code + vast entry
1eac9b8 verified
Raw History Blame Contribute Delete
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}")