File size: 1,942 Bytes
4a393d1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# tokenize rendered jsonl with the pruned tokenizer into packed arrays (ids, loss mask, doc offsets)
import json, sys, os, glob, numpy as np
from transformers import AutoTokenizer
from multiprocessing import Pool
# usage: tokenize_data.py <model_dir> <out_dir> [jsonl ...], all of data/rendered when none is given
# documents longer than MAX_TOKENS are dropped so that every document fits in one training window
MAX_TOKENS=2046
model_dir=sys.argv[1]; out_dir=sys.argv[2]; os.makedirs(out_dir,exist_ok=True)
tok=None
def init():
    global tok; tok=AutoTokenizer.from_pretrained(model_dir)
def work(lines):
    docs=[json.loads(l) for l in lines]
    enc=tok([d["text"] for d in docs],return_offsets_mapping=True,add_special_tokens=False)
    out=[]
    for d,ids,offs in zip(docs,enc["input_ids"],enc["offset_mapping"]):
        mask=np.zeros(len(ids),dtype=np.uint8)
        starts=np.array([o[0] for o in offs]); ends=np.array([o[1] for o in offs])
        for a,b in d["spans"]:
            mask[(ends>a)&(starts<b)]=1
        if len(ids)<=MAX_TOKENS: out.append((np.array(ids,dtype=np.int32),mask))
    return out
if __name__=="__main__":
    with Pool(16,initializer=init) as p:
        for f in sys.argv[3:] or sorted(glob.glob("data/rendered/*.jsonl")):
            name=os.path.basename(f)[:-6]
            lines=[l for l in open(f).read().split("\n") if l]
            chunks=[lines[i:i+512] for i in range(0,len(lines),512)]
            ids=[];masks=[];offs=[0]
            for res in p.imap(work,chunks):
                for a,m in res: ids.append(a); masks.append(m); offs.append(offs[-1]+len(a))
            ids=np.concatenate(ids); masks=np.concatenate(masks); offs=np.array(offs,dtype=np.int64)
            np.savez(os.path.join(out_dir,name+".npz"),ids=ids,mask=masks,offs=offs)
            print(f"{name}: docs {len(offs)-1} tokens {len(ids)} loss-tokens {int(masks.sum())} mean-len {len(ids)/(len(offs)-1):.0f}",flush=True)