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()