Download code/build_text2.py from fnruha0921/omnijev-work: direct link, hf CLI and curl.
- Browser
- Download file 7.08 kB
-
https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/build_text2.py
- Command line
-
hf download hf://fnruha0921/omnijev-work/code/build_text2.py
-
curl -L -o build_text2.py https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/build_text2.py
7.08 kB
| """Extra TEXT training data, disjoint from every benchmark we report (no AG News, DAIR Emotion, Banking77, MASSIVE, | |
| XNLI, JevBench, typed-decisions). Appended to DATA["text2"] (train split only).""" | |
| import random, traceback | |
| from datasets import load_dataset | |
| from mmjev import Seg | |
| rng = random.Random(1) | |
| out = [] | |
| def add(task, state, q, y): | |
| r = rec(task, "text", "train", [Seg("text", state)], q, y) | |
| r["raw"] = None | |
| out.append(r) | |
| def safe(fn): | |
| try: | |
| fn() | |
| except Exception: | |
| log("text2 skip", fn.__name__, traceback.format_exc()[-300:]) | |
| def mnli(n=900): | |
| ds = load_dataset("nyu-mll/multi_nli", split="train").shuffle(seed=1).select(range(n)) | |
| crit = {"entailment": "the premise implies the hypothesis", | |
| "neutral": "the premise neither implies nor contradicts the hypothesis", | |
| "contradiction": "the premise contradicts the hypothesis"} | |
| for i, ex in enumerate(ds): | |
| if ex["label"] not in (0, 1, 2): | |
| continue | |
| st = f"Premise: {ex['premise']}\nHypothesis: {ex['hypothesis']}" | |
| if i % 2: | |
| add("mnli", st, {"type": "choice", "instructions": "What is the relation between the premise and the hypothesis?", | |
| "criteria": crit}, ex["label"]) | |
| else: | |
| add("mnli", st, {"type": "noul", "instructions": "Does the premise imply the hypothesis?"}, int(ex["label"] == 0)) | |
| def clinc(n=700): | |
| ds = load_dataset("clinc/clinc_oos", "plus", split="train").shuffle(seed=1) | |
| names = ds.features["intent"].names | |
| ds = ds.filter(lambda e: names[e["intent"]] != "oos").select(range(n)) | |
| for ex in ds: | |
| k = rng.choice([8, 16, 32, 64, 100]) | |
| gold = names[ex["intent"]] | |
| opts = [gold] + rng.sample([x for x in names if x not in (gold, "oos")], k - 1) | |
| rng.shuffle(opts) | |
| add("clinc150", ex["text"], {"type": "choice", "instructions": "Which single label best describes the input text?", | |
| "criteria": {o.replace("_", " "): "" for o in opts}}, opts.index(gold)) | |
| def goemo(n=700): | |
| ds = load_dataset("google-research-datasets/go_emotions", "simplified", split="train").shuffle(seed=1) | |
| names = ds.features["labels"].feature.names | |
| ds = ds.filter(lambda e: len(e["labels"]) == 1 and names[e["labels"][0]] != "neutral").select(range(n)) | |
| for ex in ds: | |
| gold = names[ex["labels"][0]] | |
| k = rng.choice([6, 10, 27]) | |
| opts = [gold] + rng.sample([x for x in names if x not in (gold, "neutral")], k - 1) | |
| rng.shuffle(opts) | |
| add("goemotions", ex["text"], {"type": "choice", "instructions": "Which emotion does the text express?", | |
| "criteria": {o: "" for o in opts}}, opts.index(gold)) | |
| def yahoo(n=500): | |
| ds = load_dataset("community-datasets/yahoo_answers_topics", split="train").shuffle(seed=1).select(range(n)) | |
| names = ds.features["topic"].names | |
| for ex in ds: | |
| st = f"Question: {ex['question_title']} {ex['question_content']}\nAnswer: {ex['best_answer']}"[:1500] | |
| add("yahoo_topics", st, {"type": "choice", "instructions": "Which topic is this post about?", | |
| "criteria": {x: "" for x in names}}, ex["topic"]) | |
| def yelp(n=500): | |
| ds = load_dataset("Yelp/yelp_review_full", split="train").shuffle(seed=1).select(range(n)) | |
| lv = ["1 star, terrible", "2 stars, poor", "3 stars, average", "4 stars, good", "5 stars, excellent"] | |
| for ex in ds: | |
| add("yelp", ex["text"][:1500], {"type": "score", "instructions": "What star rating does this review give?", | |
| "criteria": lv}, ex["label"]) | |
| def policies(n=600): | |
| """JevBench-style bounded policy checks, programmatic gold.""" | |
| rules = [("a receipt is provided", "the customer has a receipt", "the customer has no receipt"), | |
| ("the purchase was made within 30 days", "the purchase was 12 days ago", "the purchase was 45 days ago"), | |
| ("the item is unused", "the item is unopened", "the item has been used"), | |
| ("the account is verified", "the account is verified", "the account is not verified"), | |
| ("the amount is below $500", "the amount is $120", "the amount is $1,300"), | |
| ("a manager approved it", "a manager approved it", "no manager approval was given"), | |
| ("the request comes from the account owner", "the account owner made the request", | |
| "a third party made the request")] | |
| actions = ["issue a refund", "approve the transfer", "grant the exception", "release the order"] | |
| for _ in range(n): | |
| rs = rng.sample(rules, rng.randint(2, 3)) | |
| ok = [rng.random() < 0.7 for _ in rs] | |
| facts = [r[1] if o else r[2] for r, o in zip(rs, ok)] | |
| rng.shuffle(facts) | |
| act = rng.choice(actions) | |
| st = (f"Policy: we {act} only if " + " and ".join(r[0] for r in rs) + ". Case: " + "; ".join(facts) + ".") | |
| if rng.random() < 0.5: | |
| add("policy_syn", st, {"type": "noul", "instructions": f"Under the stated policy, may we {act}?"}, int(all(ok))) | |
| else: | |
| n_ok = sum(ok) | |
| add("policy_syn", st, {"type": "choice", "instructions": "What should happen under the policy?", | |
| "criteria": {"allow": "every condition holds", "deny": "at least one condition fails"}}, | |
| 0 if all(ok) else 1) | |
| def snake(n=400): | |
| """Snake-game states (demo task, disclosed in the card); gold = safe move that best approaches the food, | |
| ties broken by free space (flood fill).""" | |
| import sys; sys.path.insert(0, "/content") | |
| from demos import Snake, DIRS, OPP | |
| def space(g, d): | |
| hx, hy = g.body[0]; start = (hx + DIRS[d][0], hy + DIRS[d][1]) | |
| seen, st, blocked = {start}, [start], set(g.body[:-1]) | |
| while st and len(seen) < 60: | |
| x, y = st.pop() | |
| for dx, dy in DIRS.values(): | |
| c = (x + dx, y + dy) | |
| if 0 <= c[0] < g.n and 0 <= c[1] < g.n and c not in blocked and c not in seen: | |
| seen.add(c); st.append(c) | |
| return len(seen) | |
| def best(g): | |
| hx, hy = g.body[0]; fx, fy = g.food | |
| cand = [d for d in DIRS if d != OPP[g.dir] and not g.blocked(d)] | |
| if not cand: | |
| return None | |
| return max(cand, key=lambda d: (space(g, d) >= len(g.body), -(abs(fx - hx - DIRS[d][0]) + abs(fy - hy - DIRS[d][1])), space(g, d))) | |
| made = 0 | |
| for seed in range(10_000): | |
| g = Snake(seed=seed) | |
| while g.alive and made < n and g.steps < 120: | |
| b = best(g) | |
| if b is None: | |
| break | |
| if rng.random() < 0.35: | |
| q = g.question() | |
| add("snake_syn", g.state_text(), q, list(q["criteria"]).index(b)); made += 1 | |
| g.step(b if rng.random() > 0.1 else rng.choice([d for d in DIRS if d != OPP[g.dir]])) | |
| if made >= n: | |
| break | |
| for f in (mnli, clinc, goemo, yahoo, yelp, policies, snake): | |
| safe(f) | |
| DATA["text2"] = out | |