small-test / src /prune_tokenizer.py
Serveurperso's picture
Serveurperso HF Staff
small-test: a 95M multimodal fixture for the llama.cpp server CI
4a393d1
Raw History Blame Contribute Delete
3.5 kB
import os
# build a pruned byte-level BPE tokenizer.json from token frequencies
import json, numpy as np, os, sys
from transformers import AutoTokenizer, PreTrainedTokenizerFast
P=os.environ.get("PARENT","Qwen/Qwen3.5-0.8B")
def prune_tokenizer(min_count, out_dir, freq_path="tokfreq.npy", boost=()):
t=json.load(open(os.path.join(P,"tokenizer.json")))
vocab=t["model"]["vocab"]; merges=t["model"]["merges"]
id2tok={i:s for s,i in vocab.items()}
cnt=np.load(freq_path)
# merge lookup: child -> (a,b)
def split(m): return tuple(m) if isinstance(m,list) else tuple(m.split(" ",1))
merges=[split(m) for m in merges]
parent={}
for a,b in merges: parent.setdefault(a+b,(a,b))
keep=set()
byte_toks=[s for s,i in vocab.items() if len(s)==1] # 256 base byte tokens
keep.update(byte_toks)
for i in np.nonzero(cnt>=min_count)[0]:
if i in id2tok: keep.add(id2tok[i])
for s in boost: keep.add(s)
# closure: kept token must be reachable through kept merge parents
stack=list(keep)
while stack:
s=stack.pop()
if s in parent:
for x in parent[s]:
if x not in keep: keep.add(x); stack.append(x)
# keep original order of ids for stability
kept_ids=sorted(vocab[s] for s in keep)
new_vocab={id2tok[i]:n for n,i in enumerate(kept_ids)}
new_merges=[a+" "+b for a,b in merges if a in new_vocab and b in new_vocab and (a+b) in new_vocab]
t["model"]["vocab"]=new_vocab; t["model"]["merges"]=new_merges
# added/special tokens go right after the base vocab, same order as original
added=sorted(t["added_tokens"],key=lambda a:a["id"])
old2new={i:n for n,i in enumerate(kept_ids)}
n=len(new_vocab)
for a in added:
old2new[a["id"]]=n; a["id"]=n; n+=1
os.makedirs(out_dir,exist_ok=True)
json.dump(t,open(os.path.join(out_dir,"tokenizer.json"),"w"),ensure_ascii=False)
# tokenizer_config: keep as is (special token strings unchanged)
tc=json.load(open(os.path.join(P,"tokenizer_config.json")))
tc.pop("added_tokens_decoder",None)
tc["extra_special_tokens"]={k:v for k,v in tc.get("extra_special_tokens",{}).items() if "audio" not in k}
json.dump(tc,open(os.path.join(out_dir,"tokenizer_config.json"),"w"),indent=1,ensure_ascii=False)
import shutil; shutil.copy(os.path.join(P,"chat_template.jinja"),out_dir)
print(f"pruned tokenizer: base vocab {len(new_vocab)}, merges {len(new_merges)} (from {len(merges)}), total {n}")
return old2new, n
if __name__=="__main__":
min_count=int(sys.argv[1]); out=sys.argv[2]
old2new,n=prune_tokenizer(min_count,out)
np.save(os.path.join(out,"old2new.npy"),np.array([[k,v] for k,v in old2new.items()]))
# verify round trip on samples
old=AutoTokenizer.from_pretrained(P); new=AutoTokenizer.from_pretrained(out)
import glob
same=0;tot=0;longer=0;bad=0
for f in glob.glob("data/rendered/*.eval.jsonl")+["data/rendered/smoltalk_eval.jsonl"]:
for k,l in enumerate(open(f)):
if k>=300: break
s=json.loads(l)["text"]
a=old(s)["input_ids"]; b=new(s)["input_ids"]
if new.decode(b)!=s: bad+=1
tot+=len(a); longer+=len(b)-len(a)
print(f"roundtrip mismatches: {bad}; token count +{longer/tot*100:.2f}% vs original")
print("special:", [(x, new.convert_tokens_to_ids(x)) for x in ["<|im_start|>","<|im_end|>","<tool_call>","<|endoftext|>","<|image_pad|>","<|vision_start|>","<|vision_end|>"]])