Upload tokenizer.py
Browse files- tokenizer.py +23 -0
tokenizer.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
class DigitTokenizer:
|
| 2 |
+
def __init__(self):
|
| 3 |
+
self.digit_char = "0123456789*%="
|
| 4 |
+
self.special_tokens = ["<bos>", "<pad>", "<eos>", "<unk>"]
|
| 5 |
+
self.vocab_list = list(self.digit_char) + self.special_tokens
|
| 6 |
+
|
| 7 |
+
self.char_to_int = {char: i for i, char in enumerate(self.vocab_list)}
|
| 8 |
+
self.int_to_char = {i: char for i, char in enumerate(self.vocab_list)}
|
| 9 |
+
self.vocab_size = len(self.vocab_list)
|
| 10 |
+
|
| 11 |
+
def encode(self, input_string: str) -> list[int]:
|
| 12 |
+
return [self.char_to_int.get(char, self.char_to_int["<unk>"]) for char in input_string]
|
| 13 |
+
|
| 14 |
+
def decode(self, token_ids: list[int]) -> str:
|
| 15 |
+
decoded_chars = []
|
| 16 |
+
for token_id in token_ids:
|
| 17 |
+
char = self.int_to_char.get(token_id)
|
| 18 |
+
if char == "<eos>":
|
| 19 |
+
break
|
| 20 |
+
if char is not None and char not in ["<pad>", "<unk>"]:
|
| 21 |
+
decoded_chars.append(char)
|
| 22 |
+
|
| 23 |
+
return "".join(decoded_chars)
|