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()