File size: 2,765 Bytes
cda954d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Build a runnable example from AG News, so the loop can be tried without
private data.

AG News is a public topic-classification set: a headline plus a short body,
labelled into four categories. It gives one `choice` question -- enough to see
the loop work end to end. The other two question types (`score`, `bool`) are
documented in data.py; adding them only means adding fields to the JSONL lines.

    python example/build_example.py --train 4000 --gold 400

Writes example/train.jsonl, example/gold.jsonl and example/questions.json.
The two sets are disjoint by construction.
"""
import argparse, json, os, random

LABELS = {0: "world", 1: "sports", 2: "business", 3: "science"}

QUESTIONS = {
    "topic": {
        "type": "choice",
        "instructions": "Which section of a newspaper does this story belong to?",
        "criteria": {
            "world": "international affairs, politics, conflict, society",
            "sports": "matches, athletes, tournaments, results",
            "business": "companies, markets, economy, trade",
            "science": "technology, research, health, the internet",
        },
    }
}


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--train", type=int, default=4000)
    p.add_argument("--gold", type=int, default=400)
    p.add_argument("--seed", type=int, default=0)
    a = p.parse_args()

    from datasets import load_dataset
    ds = load_dataset("fancyzhx/ag_news", split="train")
    idx = list(range(len(ds)))
    random.Random(a.seed).shuffle(idx)
    idx = idx[:a.train + a.gold]

    here = os.path.dirname(os.path.abspath(__file__))
    rows = []
    for i in idx:
        r = ds[i]
        text = r["text"].strip()
        if not text:
            continue
        rows.append({"id": "ag-%d" % i, "text": text, "topic": LABELS[r["label"]]})

    gold, train = rows[:a.gold], rows[a.gold:]
    for name, part in (("train", train), ("gold", gold)):
        path = os.path.join(here, name + ".jsonl")
        with open(path, "w") as f:
            for d in part:
                f.write(json.dumps(d, ensure_ascii=False) + "\n")
        print("%s: %d records -> %s" % (name, len(part), path))

    qpath = os.path.join(here, "questions.json")
    json.dump(QUESTIONS, open(qpath, "w"), indent=2)
    print("questions ->", qpath)

    # A number worth knowing before you read any result: the majority class.
    # A model that always answers with the most common label already scores
    # this much, and anything at or below it has learned nothing.
    from collections import Counter
    c = Counter(d["topic"] for d in gold)
    print("gold distribution:", dict(c))
    print("majority-class baseline: %.4f" % (max(c.values()) / len(gold)))


if __name__ == "__main__":
    main()