| 13 | |
| 14 | |
| 15 | class Tokenizer: |
| 16 | def __init__(self, model_path: str): |
| 17 | """ |
| 18 | Create a tokenizer, with inner implementation either spm or HF transformers tokenzier |
| 19 | :param model_path: |
| 20 | - when using spm tokenizer, should be path to a sentencepiece model with suffix `.model` |
| 21 | - when using huggingface transformers tokenizer, should be an HF model repo or a local directory, |
| 22 | containing tokenizer.json and tokenizer_config.json. |
| 23 | """ |
| 24 | if model_path.endswith(".model"): # spm tokenizer |
| 25 | self.tokenizer_type = "spm" |
| 26 | # reload tokenizer |
| 27 | assert os.path.isfile(model_path), model_path |
| 28 | self.tokenizer = SentencePieceProcessor(model_file=model_path) |
| 29 | logger.info(f"Reloaded SentencePiece model from {model_path}") |
| 30 | |
| 31 | # BOS / EOS token IDs |
| 32 | self.bos_id: int = self.tokenizer.bos_id() |
| 33 | self.eos_id: int = self.tokenizer.eos_id() |
| 34 | assert self.tokenizer.vocab_size() == self.tokenizer.get_piece_size() |
| 35 | else: |
| 36 | self.tokenizer_type = "transformers" |
| 37 | self.tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) |
| 38 | logger.info(f"load HF transformers tokenizer from {model_path}") |
| 39 | # BOS / EOS token IDs |
| 40 | self.bos_id: int = self.tokenizer.bos_token_id |
| 41 | if self.bos_id is None: |
| 42 | self.bos_id = self.tokenizer.eos_token_id |
| 43 | self.eos_id: int = self.tokenizer.eos_token_id |
| 44 | assert self.eos_id is not None |
| 45 | |
| 46 | self._probe_tokenizer_style() |
| 47 | |
| 48 | logger.info( |
| 49 | f"#words: {self.n_words} - BOS ID: {self.bos_id} - EOS ID: {self.eos_id}" |
| 50 | ) |
| 51 | |
| 52 | def encode(self, s: str, bos: bool, eos: bool) -> List[int]: |
| 53 | assert type(s) is str |
| 54 | if self.tokenizer_type == "transformers": |
| 55 | t = self.tokenizer.encode(s, truncation=False, add_special_tokens=False) |
| 56 | else: |
| 57 | t = self.tokenizer.encode(s) |
| 58 | if bos: |
| 59 | t = [self.bos_id] + t |
| 60 | if eos: |
| 61 | t = t + [self.eos_id] |
| 62 | return t |
| 63 | |
| 64 | def encode_segment(self, s:str): |
| 65 | s = s.lstrip(' ') |
| 66 | if self.need_space_before_segment: |
| 67 | return self.encode(" " + s, bos=False, eos=False) |
| 68 | else: |
| 69 | return self.encode(s, bos=False, eos=False) |
| 70 | |
| 71 | def encode_wo_prefix_space(self, s:str): |
| 72 | if self.need_space_before_segment: |