""" From-scratch byte-level BPE tokenizer, trained on the corpus (no pretrained weights). Special tokens are reserved for the hybrid model: """ import json import os from tokenizers import Tokenizer from tokenizers.models import BPE from tokenizers.trainers import BpeTrainer from tokenizers.pre_tokenizers import ByteLevel from tokenizers.decoders import ByteLevel as ByteLevelDecoder from tokenizers import AddedToken SPECIAL = ["", "", "", "", "", "", "", "", "", # learned secondary-memory tokens (side store, not in weights) "", "", "", ""] class YKTokenizer: def __init__(self, path=None): self.path = path self.tok = None self._ids = {} # ---- training -------------------------------------------------------- def train(self, iterator, vocab_size=32768, save_path=None): assert vocab_size > len(SPECIAL) model = BPE(unk_token="") self.tok = Tokenizer(model) self.tok.pre_tokenizer = ByteLevel() self.tok.decoder = ByteLevelDecoder() trainer = BpeTrainer( vocab_size=vocab_size, special_tokens=SPECIAL + [""], show_progress=True, ) self.tok.train_from_iterator(iterator, trainer) self._refresh() if save_path: self.save(save_path) return self # ---- load / save ---------------------------------------------------- def save(self, path): self.path = path os.makedirs(os.path.dirname(path), exist_ok=True) self.tok.save(path) with open(path + ".meta.json", "w") as f: json.dump({"vocab_size": self.vocab_size, "special_ids": self._ids}, f) @classmethod def load(cls, path): obj = cls(path) obj.tok = Tokenizer.from_file(path) meta = path + ".meta.json" if os.path.exists(meta): with open(meta) as f: m = json.load(f) obj._ids = m["special_ids"] obj._refresh() return obj def _refresh(self): self._ids = {s: self.tok.token_to_id(s) for s in SPECIAL} self._ids[""] = self.tok.token_to_id("") self.vocab_size = self.tok.get_vocab_size() # ---- ids ------------------------------------------------------------- @property def pad_id(self): return self._ids[""] @property def bos_id(self): return self._ids[""] @property def eos_id(self): return self._ids[""] @property def mask_id(self): return self._ids[""] @property def think_id(self): return self._ids[""] @property def endthink_id(self): return self._ids[""] @property def tool_id(self): return self._ids[""] @property def endtool_id(self): return self._ids[""] @property def result_id(self): return self._ids[""] @property def mem_write_id(self): return self._ids[""] @property def mem_read_id(self): return self._ids[""] @property def mem_kv_id(self): return self._ids[""] @property def mem_evict_id(self): return self._ids[""] # ---- encode / decode ------------------------------------------------ def encode(self, text, add_special=False): ids = self.tok.encode(text).ids if add_special: ids = [self.bos_id] + ids + [self.eos_id] return ids def encode_batch(self, texts): return [t.ids for t in self.tok.encode_batch(texts)] def decode(self, ids, skip_special=True): if skip_special: ids = [i for i in ids if i not in self._ids.values() or i == self._ids.get("", -1)] return self.tok.decode(ids, skip_special_tokens=skip_special) def __len__(self): return self.vocab_size