Download example/build_example.py from InfinimindCreations/laya-rlcd-training: direct link, hf CLI and curl.
- Browser
- Download file 2.77 kB
-
https://huggingface.co/InfinimindCreations/laya-rlcd-training/resolve/main/example/build_example.py
- Command line
-
hf download hf://InfinimindCreations/laya-rlcd-training/example/build_example.py
-
curl -L -o build_example.py https://huggingface.co/InfinimindCreations/laya-rlcd-training/resolve/main/example/build_example.py
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() | |