File size: 7,079 Bytes
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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