TinyQuery-140M / tinyquery /prepare.py
karmx's picture
Release TinyQuery 139.7M from scratch with frozen weights, reproducible Mac evaluations and runtime source
b296ad4 verified
Raw
History Blame Contribute Delete
3.43 kB
"""Learn a byte BPE using training text only; tokenize complete supervised records."""
import argparse
import hashlib
import json
from pathlib import Path
import numpy as np
from tokenizers import Tokenizer, models, trainers, pre_tokenizers, decoders
SPECIAL=['<|pad|>','<|context|>','<|user|>','<|assistant|>','<|end|>']
def file_sha256(path):
h=hashlib.sha256()
with Path(path).open('rb') as stream:
for block in iter(lambda:stream.read(4*1024*1024),b''):h.update(block)
return h.hexdigest()
def rows(path):
with open(path) as stream:
for line in stream: yield json.loads(line)
def main():
p=argparse.ArgumentParser(); p.add_argument('--data',default='data/tinyquery')
p.add_argument('--vocab',type=int,default=24000); p.add_argument('--context',type=int,default=2048)
p.add_argument('--reuse-tokenizer',action='store_true')
args=p.parse_args(); base=Path(args.data)
if args.reuse_tokenizer:
tokenizer=Tokenizer.from_file(str(base/'tokenizer.json'))
else:
tokenizer=Tokenizer(models.BPE())
tokenizer.pre_tokenizer=pre_tokenizers.ByteLevel(add_prefix_space=False)
tokenizer.decoder=decoders.ByteLevel()
trainer=trainers.BpeTrainer(vocab_size=args.vocab,min_frequency=2,special_tokens=SPECIAL,
initial_alphabet=pre_tokenizers.ByteLevel.alphabet())
tokenizer.train_from_iterator((r['prompt']+r['response']+'<|end|>' for r in rows(base/'train.jsonl')),trainer)
tokenizer.save(str(base/'tokenizer.json'))
stats={'vocab_size':tokenizer.get_vocab_size(),'context':args.context,'special_tokens':{s:tokenizer.token_to_id(s) for s in SPECIAL}}
for split in ['train','validation','test']:
data=list(rows(base/(split+'.jsonl')))
encoded=[]; boundaries=[]; lengths=[]; actions=[]; kept=[]; rejected=[]
for i,r in enumerate(data):
prefix=tokenizer.encode(r['prompt']).ids
suffix=tokenizer.encode(r['response']+'<|end|>').ids
ids=prefix+suffix
if len(ids)>args.context+1:
rejected.append({'id':r['id'],'tokens':len(ids)}); continue
encoded.append(ids); boundaries.append(len(prefix)-1); lengths.append(len(ids))
actions.append({'call':0,'clarify':1,'answer':2}[r['target']['action']]); kept.append(i)
width=max(lengths)
array=np.zeros((len(encoded),width),dtype=np.uint16)
for i,ids in enumerate(encoded): array[i,:len(ids)]=ids
np.save(base/(split+'-ids.npy'),array)
np.savez(base/(split+'-meta.npz'),boundaries=np.array(boundaries,dtype=np.int32),
lengths=np.array(lengths,dtype=np.int32),actions=np.array(actions,dtype=np.int32),
kept=np.array(kept,dtype=np.int32),sample_weights=np.array([data[i].get('sample_weight',1.0) for i in kept],dtype=np.float32))
stats[split]={'examples':len(encoded),'tokens':sum(lengths),'response_tokens':sum(l-b-1 for l,b in zip(lengths,boundaries)),
'ids_sha256':file_sha256(base/(split+'-ids.npy')),
'length_p50':float(np.median(lengths)),'length_p95':float(np.quantile(lengths,.95)),
'max_length':max(lengths),'rejected':rejected}
print(split,stats[split],flush=True)
(base/'tokenization.json').write_text(json.dumps(stats,indent=2))
if __name__=='__main__': main()