omnijev-work / code /build_text2.py
fnruha0921's picture
code for H200 job
8c867d9 verified
Raw History Blame Contribute Delete
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