| 6 | |
| 7 | |
| 8 | class Tokenizer: |
| 9 | def __init__(self, model_name: str): |
| 10 | self._tokenizer = transformers.AutoTokenizer.from_pretrained( |
| 11 | model_name, |
| 12 | legacy=False, |
| 13 | model_max_length=512, |
| 14 | ) |
| 15 | self._decoder_start_id = 0 |
| 16 | |
| 17 | @property |
| 18 | def eos_id(self) -> int: |
| 19 | return self._tokenizer.eos_token_id |
| 20 | |
| 21 | @property |
| 22 | def decoder_start_id(self) -> int: |
| 23 | return self._decoder_start_id |
| 24 | |
| 25 | def encode(self, s: str) -> mx.array: |
| 26 | return mx.array( |
| 27 | self._tokenizer( |
| 28 | s, |
| 29 | return_tensors="np", |
| 30 | return_attention_mask=False, |
| 31 | )[ |
| 32 | "input_ids" |
| 33 | ].squeeze(0) |
| 34 | ) |
| 35 | |
| 36 | def decode(self, t: List[int]) -> str: |
| 37 | return self._tokenizer.decode(t) |
| 38 | |
| 39 | |
| 40 | class SpeculativeDecoder: |