FST_code / src /lmr /tokenizer /tokenizer_bert.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
5.34 kB
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]