"""Tokenize raw shards into uint16 .bin files, one per input file, docs separated by <|endoftext|>. Every 200th document goes to a validation file instead. Re-running skips finished files. Usage: source env.sh && $TA_PY scripts/tokenize_data.py [--workers 24] Output: $TA_DATA/tok/{train,val}/{source}/{file}.bin and $TA_DATA/tok/manifest.json """ import argparse import json import os from multiprocessing import Pool os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") import numpy as np from tokenizers import Tokenizer from tiny_agent.text import DATA, EOS_ID, source_files BATCH = 512 VAL_EVERY = 200 def work(args): source, path, it_name = args from tiny_agent import text it = getattr(text, it_name) base = os.path.basename(path).split(".")[0] out_tr = f"{DATA}/tok/train/{source}/{base}.bin" out_va = f"{DATA}/tok/val/{source}/{base}.bin" if os.path.exists(out_tr) and os.path.exists(out_va): return source, path, os.path.getsize(out_tr) // 2, os.path.getsize(out_va) // 2, "cached" os.makedirs(os.path.dirname(out_tr), exist_ok=True) os.makedirs(os.path.dirname(out_va), exist_ok=True) tok = Tokenizer.from_file(f"{DATA}/tokenizer.json") n_tr = n_va = 0 with open(out_tr + ".tmp", "wb") as ftr, open(out_va + ".tmp", "wb") as fva: batch, k = [], 0 def flush(batch, k): nonlocal n_tr, n_va for enc in tok.encode_batch(batch, add_special_tokens=False): arr = np.asarray(enc.ids + [EOS_ID], dtype=np.uint16) if k % VAL_EVERY == VAL_EVERY - 1: arr.tofile(fva) n_va += arr.size else: arr.tofile(ftr) n_tr += arr.size k += 1 return k for t in it(path): batch.append(t) if len(batch) == BATCH: k = flush(batch, k) batch = [] if batch: flush(batch, k) os.replace(out_tr + ".tmp", out_tr) os.replace(out_va + ".tmp", out_va) return source, path, n_tr, n_va, "done" def main(): ap = argparse.ArgumentParser() ap.add_argument("--workers", type=int, default=24) a = ap.parse_args() jobs = [(s, f, it.__name__) for s, files in source_files().items() for f, it in files] manifest = {} with Pool(a.workers) as pool: for source, path, n_tr, n_va, status in pool.imap_unordered(work, jobs): m = manifest.setdefault(source, {"train_tokens": 0, "val_tokens": 0, "files": 0}) m["train_tokens"] += n_tr m["val_tokens"] += n_va m["files"] += 1 print(f"{status} {source} {os.path.basename(path)} train={n_tr/1e6:.1f}M val={n_va/1e6:.2f}M", flush=True) tot = sum(m["train_tokens"] for m in manifest.values()) manifest["_total_train_tokens"] = tot with open(f"{DATA}/tok/manifest.json", "w") as f: json.dump(manifest, f, indent=1) print(json.dumps(manifest, indent=1)) if __name__ == "__main__": main()