| """Small, auditable byte-pair encoder trained only on supplied local text.""" |
| import argparse, collections, json |
| from pathlib import Path |
| SPECIAL = ["<pad>", "<bos>", "<eos>", "<system>", "<user>", "<assistant>", "<tool>", "<plan>", "<memory>"] |
|
|
| class ByteBPETokenizer: |
| def __init__(self, merges=None, special_tokens=SPECIAL): |
| self.special_tokens = special_tokens |
| self.special_to_id = {s:i for i,s in enumerate(special_tokens)} |
| self.merges = [tuple(x) for x in (merges or [])] |
| self.base = len(special_tokens) |
| self.pair_to_token = {p:self.base+256+i for i,p in enumerate(self.merges)} |
| self.pair_rank = {p:i for i,p in enumerate(self.merges)} |
| self.vocab_size = self.base + 256 + len(self.merges) |
| def encode_bytes(self, raw): |
| ids = [self.base+b for b in raw] |
| |
| while len(ids) > 1: |
| choices=[(i, self.pair_rank[(ids[i],ids[i+1])]) for i in range(len(ids)-1) if (ids[i],ids[i+1]) in self.pair_rank] |
| if not choices: break |
| i,rank=min(choices,key=lambda x:x[1]); ids[i:i+2]=[self.base+256+rank] |
| return ids |
| def encode(self, text, add_bos=False, add_eos=False): |
| ids=[]; pos=0 |
| while pos < len(text): |
| found=next(((s,i) for s,i in self.special_to_id.items() if text.startswith(s,pos)),None) |
| if found: ids.append(found[1]);pos+=len(found[0]);continue |
| ends=[text.find(s,pos) for s in self.special_to_id if text.find(s,pos)>=0]; end=min(ends) if ends else len(text) |
| ids.extend(self.encode_bytes(text[pos:end].encode('utf8')));pos=end |
| return ([self.special_to_id['<bos>']] if add_bos else [])+ids+([self.special_to_id['<eos>']] if add_eos else []) |
| def decode(self, ids): |
| def expand(x): |
| if x < self.base: return b'' |
| if x < self.base+256: return bytes([x-self.base]) |
| a,b=self.merges[x-self.base-256]; return expand(a)+expand(b) |
| out=[];data=bytearray() |
| for x in ids: |
| if x < self.base: |
| if data: out.append(data.decode('utf8',errors='replace'));data.clear() |
| out.append(self.special_tokens[x]) |
| else:data.extend(expand(x)) |
| if data:out.append(data.decode('utf8',errors='replace')) |
| return ''.join(out) |
| def save(self,path): Path(path).parent.mkdir(parents=True,exist_ok=True);Path(path).write_text(json.dumps({'merges':self.merges,'special_tokens':self.special_tokens})) |
| @classmethod |
| def load(cls,path): |
| d=json.loads(Path(path).read_text());return cls(d['merges'],d['special_tokens']) |
|
|
| def train(files,vocab_size): |
| symbols=[len(SPECIAL)+b for b in b''.join(Path(p).read_bytes() for p in files)];merges=[] |
| while len(SPECIAL)+256+len(merges)<vocab_size: |
| pairs=collections.Counter(zip(symbols,symbols[1:])); |
| if not pairs:break |
| pair,n=pairs.most_common(1)[0] |
| if n<2:break |
| token=len(SPECIAL)+256+len(merges);merges.append(pair);new=[];i=0 |
| while i<len(symbols): |
| if i+1<len(symbols) and (symbols[i],symbols[i+1])==pair:new.append(token);i+=2 |
| else:new.append(symbols[i]);i+=1 |
| symbols=new |
| return ByteBPETokenizer(merges) |
|
|
| def main(): |
| p=argparse.ArgumentParser(); sub=p.add_subparsers(dest='cmd',required=True); t=sub.add_parser('train'); t.add_argument('--input',required=True);t.add_argument('--out',required=True);t.add_argument('--vocab-size',type=int,default=16000);a=p.parse_args() |
| files=list(Path(a.input).rglob('*.txt')); assert files, 'No .txt files found'; tok=train(files,a.vocab_size);tok.save(a.out);print(f'Saved {tok.vocab_size} tokens to {a.out}') |
| if __name__=='__main__': main() |
|
|