File size: 4,071 Bytes
df43f42
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b5b2dd7
 
 
df43f42
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b5b2dd7
 
 
 
 
 
 
 
 
df43f42
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
"""
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