Gala-598M-MLX / data.py
junafinity's picture
Gala-598M: weights, code, logs, report
76bbe95 verified
Raw History Blame Contribute Delete
3.16 kB
"""
Tokenise a dataset into flat uint16 token files (GPT-2 BPE, like nanoGPT).
python data.py --dataset shakespeare # ~300K tokens, seconds
python data.py --dataset fineweb --tokens 1e9 # streams FineWeb-Edu, ~1B tokens
python data.py --dataset synthetic --tokens 5e6 # no network needed (smoke tests)
Outputs data/<name>/train.bin and data/<name>/val.bin
"""
import argparse, os, sys
import numpy as np
EOT = 50256
def write(tokens: np.ndarray, out_dir: str, val_frac: float = 0.02):
os.makedirs(out_dir, exist_ok=True)
n_val = max(1024, int(len(tokens) * val_frac))
tokens[: -n_val].astype(np.uint16).tofile(os.path.join(out_dir, "train.bin"))
tokens[-n_val:].astype(np.uint16).tofile(os.path.join(out_dir, "val.bin"))
print(f"wrote {len(tokens) - n_val:,} train / {n_val:,} val tokens to {out_dir}")
def shakespeare(out_dir):
import urllib.request, tiktoken
url = "https://raw.githubusercontent.com/karpathy/char-rnn/master/data/tinyshakespeare/input.txt"
text = urllib.request.urlopen(url).read().decode()
enc = tiktoken.get_encoding("gpt2")
toks = np.array(enc.encode_ordinary(text) + [EOT], dtype=np.uint16)
write(toks, out_dir)
def fineweb(out_dir, n_tokens: int, name="sample-10BT"):
from datasets import load_dataset
import tiktoken
enc = tiktoken.get_encoding("gpt2")
ds = load_dataset("HuggingFaceFW/fineweb-edu", name=name, split="train", streaming=True)
buf, total = [], 0
for ex in ds:
t = enc.encode_ordinary(ex["text"]) + [EOT]
buf.append(np.array(t, dtype=np.uint16))
total += len(t)
if total % 10_000_000 < len(t):
print(f" {total/1e6:.0f}M tokens", file=sys.stderr)
if total >= n_tokens:
break
write(np.concatenate(buf), out_dir)
def synthetic(out_dir, n_tokens: int, vocab: int = 4096, seed: int = 0):
"""A structured pseudo-language: a sparse Markov chain with copy patterns,
so a model can actually lower its loss. For pipeline smoke tests only."""
rng = np.random.default_rng(seed)
n_states = 64
trans = rng.dirichlet(np.ones(n_states) * 0.1, size=n_states)
emit = np.stack([rng.choice(vocab, size=32, replace=False) for _ in range(n_states)])
toks = np.empty(n_tokens, dtype=np.uint16)
s = 0
for i in range(n_tokens):
if i > 256 and rng.random() < 0.05: # copy a recent token (rewards recall)
toks[i] = toks[i - rng.integers(1, 256)]
else:
toks[i] = emit[s, rng.integers(32)]
s = rng.choice(n_states, p=trans[s])
write(toks, out_dir)
if __name__ == "__main__":
ap = argparse.ArgumentParser()
ap.add_argument("--dataset", choices=["shakespeare", "fineweb", "synthetic"], required=True)
ap.add_argument("--tokens", type=float, default=1e9)
ap.add_argument("--out", default=None)
a = ap.parse_args()
out = a.out or os.path.join("data", a.dataset)
if a.dataset == "shakespeare":
shakespeare(out)
elif a.dataset == "fineweb":
fineweb(out, int(a.tokens))
else:
synthetic(out, int(a.tokens))