From 17a81b9c86cfedd70812a0e83f33798b64c1678e Mon Sep 17 00:00:00 2001 From: yurenh Date: Mon, 31 Aug 2026 18:16:31 -0500 Subject: data prep (FineWeb-Edu->GPT2 BPE), rho diagnostics, tests, measured-cost notes Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe --- scripts/prepare_data.py | 35 +++++++++++++++++++++++++++++++++++ 1 file changed, 35 insertions(+) create mode 100644 scripts/prepare_data.py (limited to 'scripts') diff --git a/scripts/prepare_data.py b/scripts/prepare_data.py new file mode 100644 index 0000000..626e7a8 --- /dev/null +++ b/scripts/prepare_data.py @@ -0,0 +1,35 @@ +"""FineWeb-Edu -> GPT-2 BPE uint16 shards (train.bin / val.bin). + python scripts/prepare_data.py --tokens 3e9 --out data/fineweb [--dataset HuggingFaceFW/fineweb-edu --name sample-10BT]""" +import os, sys, argparse +import numpy as np + +p = argparse.ArgumentParser() +p.add_argument("--dataset", default="HuggingFaceFW/fineweb-edu") +p.add_argument("--name", default="sample-10BT") +p.add_argument("--tokens", type=float, default=3e9) +p.add_argument("--val_tokens", type=float, default=5e6) +p.add_argument("--out", default="data/fineweb") +a = p.parse_args() +os.makedirs(a.out, exist_ok=True) +import tiktoken +from datasets import load_dataset +enc = tiktoken.get_encoding("gpt2") +ds = load_dataset(a.dataset, name=a.name, split="train", streaming=True) +train_path, val_path = os.path.join(a.out, "train.bin"), os.path.join(a.out, "val.bin") +ftr, fva = open(train_path, "wb"), open(val_path, "wb") +n_tr = n_va = 0 +target_tr, target_va = int(a.tokens), int(a.val_tokens) +buf = [] +for i, ex in enumerate(ds): + ids = enc.encode_ordinary(ex["text"]) + [enc.eot_token] + arr = np.array(ids, dtype=np.uint16) + if n_va < target_va and i % 100 == 0: # every 100th doc to val until filled + fva.write(arr.tobytes()); n_va += len(arr) + else: + ftr.write(arr.tobytes()); n_tr += len(arr) + if n_tr % 50_000_000 < len(arr): + print(f"train {n_tr/1e6:.0f}M val {n_va/1e6:.1f}M tokens", flush=True) + if n_tr >= target_tr and n_va >= target_va: + break +ftr.close(); fva.close() +print(f"DONE train {n_tr/1e6:.1f}M val {n_va/1e6:.1f}M -> {a.out}") -- cgit v1.2.3