File size: 1,629 Bytes
24ebd71
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import argparse
import json
from pathlib import Path

from datasets import load_dataset

PROJECT_DIR = Path(__file__).resolve().parent
DATA_DIR = PROJECT_DIR / "data"


def write_split(split: str, limit: int, destination: Path) -> int:
    dataset = load_dataset(
        "roneneldan/TinyStories",
        split=split,
        streaming=True,
    )
    written = 0
    with destination.open("w", encoding="utf-8") as handle:
        for row in dataset:
            text = str(row.get("text", "")).strip()
            if not text:
                continue
            handle.write(json.dumps({"text": text}, ensure_ascii=False) + "\n")
            written += 1
            if written >= limit:
                break
    return written


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--train-stories", type=int, default=6000)
    parser.add_argument("--eval-stories", type=int, default=600)
    args = parser.parse_args()

    DATA_DIR.mkdir(parents=True, exist_ok=True)
    train_count = write_split("train", args.train_stories, DATA_DIR / "train.jsonl")
    eval_count = write_split("validation", args.eval_stories, DATA_DIR / "eval.jsonl")
    manifest = {
        "source": "roneneldan/TinyStories",
        "train_stories": train_count,
        "eval_stories": eval_count,
        "license_note": "See the source dataset card for dataset terms.",
    }
    (DATA_DIR / "manifest.json").write_text(
        json.dumps(manifest, indent=2),
        encoding="utf-8",
    )
    print(json.dumps(manifest, indent=2))


if __name__ == "__main__":
    main()