laya-rlcd-training / example /build_example.py
Cytrex's picture
Upload folder using huggingface_hub
cda954d verified
Raw History Blame Contribute Delete
2.77 kB
"""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()