File size: 3,073 Bytes
4397e12 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 | """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()
|