clankerDiffusion-base / tokenizer.py
coderofpears's picture
Upload tokenizer.py with huggingface_hub
b5b2dd7 verified
Raw
History Blame Contribute Delete
4.07 kB
"""
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>",
# learned secondary-memory tokens (side store, not in weights)
"<mem_write>", "<mem_read>", "<mem_kv>", "<mem_evict>"]
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="<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
# ---- 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["<unk>"] = self.tok.token_to_id("<unk>")
self.vocab_size = self.tok.get_vocab_size()
# ---- ids -------------------------------------------------------------
@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>"]
# ---- 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("<unk>", -1)]
return self.tok.decode(ids, skip_special_tokens=skip_special)
def __len__(self):
return self.vocab_size