tiny-agent-112m / code /scripts /tokenize_data.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
3.07 kB
"""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()