| import re
|
| from typing import List, Union, Optional, Dict
|
|
|
|
|
| try:
|
| import tiktoken
|
| HAVE_TIKTOKEN = True
|
| except Exception:
|
| tiktoken = None
|
| HAVE_TIKTOKEN = False
|
|
|
|
|
| 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")
|
|
|
|
|
|
|
|
|
| if self.use_hf:
|
| self.hf = AutoTokenizer.from_pretrained(
|
| base_type,
|
| use_fast=True
|
| )
|
|
|
|
|
| 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]"
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
| 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|>"]
|
|
|
|
|
| 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,
|
| )
|
|
|
|
|
|
|
|
|
| @staticmethod
|
| def get_instance():
|
| if Tokenizer._instance is None:
|
| Tokenizer._instance = Tokenizer()
|
| return Tokenizer._instance
|
|
|
| @staticmethod
|
| def set_instance(tokenizer):
|
| Tokenizer._instance = tokenizer
|
|
|
|
|
|
|
|
|
| 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")
|
|
|
|
|
|
|
|
|
| @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]
|
|
|