| """ |
| From-scratch byte-level BPE tokenizer, trained on the corpus (no pretrained weights). |
| Special tokens are reserved for the hybrid model: |
| <pad> <bos> <eos> <mask> <think> </think> <tool> </tool> <result> |
| """ |
| 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 = ["<pad>", "<bos>", "<eos>", "<mask>", |
| "<think>", "</think>", "<tool>", "</tool>", "<result>", |
| |
| "<mem_write>", "<mem_read>", "<mem_kv>", "<mem_evict>"] |
|
|
|
|
| class YKTokenizer: |
| def __init__(self, path=None): |
| self.path = path |
| self.tok = None |
| self._ids = {} |
|
|
| |
| def train(self, iterator, vocab_size=32768, save_path=None): |
| assert vocab_size > len(SPECIAL) |
| model = BPE(unk_token="<unk>") |
| self.tok = Tokenizer(model) |
| self.tok.pre_tokenizer = ByteLevel() |
| self.tok.decoder = ByteLevelDecoder() |
| trainer = BpeTrainer( |
| vocab_size=vocab_size, |
| special_tokens=SPECIAL + ["<unk>"], |
| show_progress=True, |
| ) |
| self.tok.train_from_iterator(iterator, trainer) |
| self._refresh() |
| if save_path: |
| self.save(save_path) |
| return self |
|
|
| |
| 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["<unk>"] = self.tok.token_to_id("<unk>") |
| self.vocab_size = self.tok.get_vocab_size() |
|
|
| |
| @property |
| def pad_id(self): return self._ids["<pad>"] |
| @property |
| def bos_id(self): return self._ids["<bos>"] |
| @property |
| def eos_id(self): return self._ids["<eos>"] |
| @property |
| def mask_id(self): return self._ids["<mask>"] |
| @property |
| def think_id(self): return self._ids["<think>"] |
| @property |
| def endthink_id(self): return self._ids["</think>"] |
| @property |
| def tool_id(self): return self._ids["<tool>"] |
| @property |
| def endtool_id(self): return self._ids["</tool>"] |
| @property |
| def result_id(self): return self._ids["<result>"] |
| @property |
| def mem_write_id(self): return self._ids["<mem_write>"] |
| @property |
| def mem_read_id(self): return self._ids["<mem_read>"] |
| @property |
| def mem_kv_id(self): return self._ids["<mem_kv>"] |
| @property |
| def mem_evict_id(self): return self._ids["<mem_evict>"] |
|
|
| |
| 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("<unk>", -1)] |
| return self.tok.decode(ids, skip_special_tokens=skip_special) |
|
|
| def __len__(self): |
| return self.vocab_size |
|
|