| 5 | |
| 6 | |
| 7 | class Tokenizer(object): |
| 8 | def __init__(self, provider: str, model_name: str) -> None: |
| 9 | if provider == "openai": |
| 10 | self.tokenizer = tiktoken.encoding_for_model(model_name) |
| 11 | elif provider == "huggingface": |
| 12 | self.tokenizer = LlamaTokenizer.from_pretrained(model_name) |
| 13 | # turn off adding special tokens automatically |
| 14 | self.tokenizer.add_special_tokens = False # type: ignore[attr-defined] |
| 15 | self.tokenizer.add_bos_token = False # type: ignore[attr-defined] |
| 16 | self.tokenizer.add_eos_token = False # type: ignore[attr-defined] |
| 17 | elif provider == "ours": |
| 18 | self.tokenizer = tiktoken.encoding_for_model("gpt-4") |
| 19 | else: |
| 20 | raise NotImplementedError |
| 21 | |
| 22 | def encode(self, text: str) -> list[int]: |
| 23 | return self.tokenizer.encode(text) |
| 24 | |
| 25 | def decode(self, ids: list[int]) -> str: |
| 26 | return self.tokenizer.decode(ids) |
| 27 | |
| 28 | def __call__(self, text: str) -> list[int]: |
| 29 | return self.tokenizer.encode(text) |