"""Train a 32k byte-level BPE on a sample of the actual mix; also write the Engram compressed-id map. Usage: source env.sh && $TA_PY scripts/train_tokenizer.py Outputs: $TA_DATA/tokenizer.json, $TA_DATA/cid_map.npy """ import argparse import itertools import unicodedata import numpy as np from tokenizers import Regex, Tokenizer, decoders, models, normalizers, pre_tokenizers, trainers from tiny_agent.text import DATA, SPECIAL_TOKENS, source_files # GPT-4 style split (contractions, letters, 1 digit at a time, punctuation runs, newline runs). SPLIT = (r"""'(?i:[sdmt]|ll|ve|re)|[^\r\n\p{L}\p{N}]?+\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]++[\r\n]*|\s*[\r\n]|\s+(?!\S)|\s+""") # MB of text sampled per source for training the tokenizer SAMPLE_MB = {"english": 120, "math": 90, "code_python": 70, "code_shell": 25, "code_markdown": 20, "code_javascript": 15, "code_typescript": 15, "code_go": 10, "code_rust": 10, "code_c": 10} def sample_texts(): srcs = source_files() for name, mb in SAMPLE_MB.items(): budget, got = mb * 2**20, 0 files = srcs.get(name, []) if not files: print("missing source", name, flush=True) continue f, it = files[0] for t in it(f): t = t[:20000] got += len(t) yield t if got >= budget: break print(name, got // 2**20, "MB", flush=True) def build_tokenizer() -> Tokenizer: tok = Tokenizer(models.BPE(byte_fallback=False)) tok.normalizer = normalizers.NFC() tok.pre_tokenizer = pre_tokenizers.Sequence([ pre_tokenizers.Split(Regex(SPLIT), behavior="isolated"), pre_tokenizers.ByteLevel(add_prefix_space=False, use_regex=False), ]) tok.decoder = decoders.ByteLevel() return tok def compressed_map(tok: Tokenizer) -> np.ndarray: """Engram tokenizer compression: tokens that normalize to the same surface (NFKC, accents stripped, lowercased, whitespace stripped) share one compressed id.""" V = tok.get_vocab_size() keys = {} out = np.zeros(V, dtype=np.int32) for i in range(V): s = tok.decode([i], skip_special_tokens=False) if i >= len(SPECIAL_TOKENS): s = unicodedata.normalize("NFKC", s) s = "".join(c for c in unicodedata.normalize("NFD", s) if unicodedata.category(c) != "Mn") s = s.lower().strip() or "" out[i] = keys.setdefault(s, len(keys)) print(f"compressed vocab {len(keys)} / {V} ({1 - len(keys) / V:.1%} smaller)") return out def main(): ap = argparse.ArgumentParser() ap.add_argument("--vocab", type=int, default=32768) a = ap.parse_args() tok = build_tokenizer() trainer = trainers.BpeTrainer(vocab_size=a.vocab, special_tokens=SPECIAL_TOKENS, min_frequency=2, initial_alphabet=pre_tokenizers.ByteLevel.alphabet(), show_progress=False) tok.train_from_iterator(sample_texts(), trainer=trainer) assert tok.token_to_id("<|endoftext|>") == 0 and tok.token_to_id("") == 6 tok.save(f"{DATA}/tokenizer.json") np.save(f"{DATA}/cid_map.npy", compressed_map(tok)) for s in ["def add(a, b):\n return a + b\n", "The answer is 12345.", "grep -rn 'foo' src/ | head"]: ids = tok.encode(s).ids print(len(s), "chars ->", len(ids), "tokens:", [tok.decode([i]) for i in ids][:16]) if __name__ == "__main__": main()