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]
|