File size: 5,337 Bytes
3b2d368
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
import re
from typing import List, Union, Optional, Dict

# ---- Optional tiktoken ----
try:
    import tiktoken
    HAVE_TIKTOKEN = True
except Exception:
    tiktoken = None
    HAVE_TIKTOKEN = False

# ---- HuggingFace ----
from transformers import AutoTokenizer


class Tokenizer:
    """

    Unified Tokenizer wrapper.



    - If base_type == "bert-base-uncased":

        -> HuggingFace WordPiece tokenizer (CORRECT for BERT + MLM)

    - Else:

        -> tiktoken-based tokenizer (GPT-style)

    """

    _instance = None

    def __init__(self, base_type: str = "bert-base-uncased"):
        self.base_type = base_type
        self.use_hf = (base_type == "bert-base-uncased")

        # =========================
        # BERT / WordPiece branch
        # =========================
        if self.use_hf:
            self.hf = AutoTokenizer.from_pretrained(
                base_type,
                use_fast=True
            )

            # Ensure required tokens exist
            assert self.hf.mask_token_id is not None, "BERT tokenizer must have [MASK]"
            assert self.hf.pad_token_id is not None, "BERT tokenizer must have [PAD]"

            # Expose common attributes
            self.pad_token_id = self.hf.pad_token_id
            self.eos_token_id = self.hf.sep_token_id
            self.bos_token_id = self.hf.cls_token_id
            self.cls_token_id = self.hf.cls_token_id
            self.mask_token_id = self.hf.mask_token_id
            print(self.bos_token_id,'self.hf.cls_token_id')
            return

        # =========================
        # GPT / tiktoken branch
        # =========================
        if not HAVE_TIKTOKEN:
            raise RuntimeError("tiktoken is not available, but base_type is not bert-base-uncased")

        base_enc = tiktoken.get_encoding(base_type)

        base_n = base_enc.n_vocab

        special_tokens = {
            "<|pad|>": base_n,
            "<|bos|>": base_n + 1,
            "<|eos|>": base_n + 2,
            "<|mask|>": base_n + 3,
        }

        self.special_tokens = special_tokens
        self.pad_token_id = special_tokens["<|pad|>"]
        self.bos_token_id = special_tokens["<|bos|>"]
        self.eos_token_id = special_tokens["<|eos|>"]
        self.mask_token_id = special_tokens["<|mask|>"]

        # Regex: numbers first, then default BPE
        number_pattern = r"\d{1,3}"
        pat_str = getattr(base_enc, "_pat_str", base_enc.pat_str)
        custom_pat = f"{number_pattern}|{pat_str}"

        self.tk = tiktoken.Encoding(
            name=f"{base_type}_custom",
            pat_str=custom_pat,
            mergeable_ranks=base_enc._mergeable_ranks,
            special_tokens=special_tokens,
        )

    # ------------------------------------------------------------------
    # Singleton helpers
    # ------------------------------------------------------------------
    @staticmethod
    def get_instance():
        if Tokenizer._instance is None:
            Tokenizer._instance = Tokenizer()
        return Tokenizer._instance

    @staticmethod
    def set_instance(tokenizer):
        Tokenizer._instance = tokenizer

    # ------------------------------------------------------------------
    # Encode / Decode
    # ------------------------------------------------------------------
    def encode(self, text: str, add_special_tokens: bool = True) -> List[int]:
        if self.use_hf:
            return self.hf.encode(text, add_special_tokens=add_special_tokens)
        return self.tk.encode(text, disallowed_special=())

    def decode(self, tokens: Union[List[int], List]) -> str:
        if self.use_hf:
            return self.hf.decode(tokens, skip_special_tokens=True)
        return self.tk.decode(tokens)

    def __call__(self, *args, **kwargs):
        if self.use_hf:
            return self.hf(*args, **kwargs)
        raise NotImplementedError("Call-style not supported for tiktoken branch")

    # ------------------------------------------------------------------
    # Vocab / Special tokens
    # ------------------------------------------------------------------
    @property
    def vocab_size(self) -> int:
        if self.use_hf:
            return self.hf.vocab_size
        return self.tk.n_vocab

    def get_vocab(self):
        if self.use_hf:
            return self.hf.get_vocab()
        return None

    def is_special(self, token_id: int) -> bool:
        if self.use_hf:
            return token_id in {
                self.pad_token_id,
                self.bos_token_id,
                self.eos_token_id,
                self.mask_token_id,
            }
        return token_id in self.special_tokens.values()

    def get_special_tokens_mask(self, token_ids: List[int], already_has_special_tokens=True):
        if self.use_hf:
            return self.hf.get_special_tokens_mask(
                token_ids,
                already_has_special_tokens=already_has_special_tokens
            )
        return [1 if self.is_special(t) else 0 for t in token_ids]

    def clean_tokens(self, tokens: List[int]) -> List[int]:
        return [t for t in tokens if t != self.pad_token_id and t != self.eos_token_id]