| 121 | |
| 122 | |
| 123 | class TokenizerWrapper: |
| 124 | def __init__(self, tokenizer_type: Literal["tiktoken", "huggingface"] = "tiktoken", model_name: str = "gpt-4o"): |
| 125 | self.tokenizer_type = tokenizer_type |
| 126 | self.model_name = model_name |
| 127 | self._tokenizer = None |
| 128 | self._lazy_load_tokenizer() |
| 129 | |
| 130 | def _lazy_load_tokenizer(self): |
| 131 | if self._tokenizer is not None: |
| 132 | return |
| 133 | logger.info(f"Loading tokenizer: type='{self.tokenizer_type}', name='{self.model_name}'") |
| 134 | if self.tokenizer_type == "tiktoken": |
| 135 | self._tokenizer = tiktoken.encoding_for_model(self.model_name) |
| 136 | elif self.tokenizer_type == "huggingface": |
| 137 | if AutoTokenizer is None: |
| 138 | raise ImportError("`transformers` is not installed. Please install it via `pip install transformers` to use HuggingFace tokenizers.") |
| 139 | self._tokenizer = AutoTokenizer.from_pretrained(self.model_name, use_fast=True) |
| 140 | else: |
| 141 | raise ValueError(f"Unknown tokenizer_type: {self.tokenizer_type}") |
| 142 | |
| 143 | def get_tokenizer(self): |
| 144 | """提供对底层 tokenizer 对象的访问,用于特殊情况(如 decode_batch)。""" |
| 145 | self._lazy_load_tokenizer() |
| 146 | return self._tokenizer |
| 147 | |
| 148 | def encode(self, text: str) -> list[int]: |
| 149 | self._lazy_load_tokenizer() |
| 150 | return self._tokenizer.encode(text) |
| 151 | |
| 152 | def decode(self, tokens: list[int]) -> str: |
| 153 | self._lazy_load_tokenizer() |
| 154 | return self._tokenizer.decode(tokens) |
| 155 | |
| 156 | # +++ 新增 +++: 增加一个批量解码的方法以提高效率,并保持接口一致性 |
| 157 | def decode_batch(self, tokens_list: list[list[int]]) -> list[str]: |
| 158 | self._lazy_load_tokenizer() |
| 159 | # HuggingFace tokenizer 有 decode_batch,但 tiktoken 没有,我们用列表推导来模拟 |
| 160 | if self.tokenizer_type == "tiktoken": |
| 161 | return [self._tokenizer.decode(tokens) for tokens in tokens_list] |
| 162 | elif self.tokenizer_type == "huggingface": |
| 163 | return self._tokenizer.batch_decode(tokens_list, skip_special_tokens=True) |
| 164 | else: |
| 165 | raise ValueError(f"Unknown tokenizer_type: {self.tokenizer_type}") |
| 166 | |
| 167 | |
| 168 |
no outgoing calls
no test coverage detected