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|>","","<|endoftext|>","<|image_pad|>","<|vision_start|>","<|vision_end|>"]])