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]