| 8 | |
| 9 | class TestBpe(unittest.TestCase): |
| 10 | def test_bpe(self): |
| 11 | symbols = CharTrie() |
| 12 | symbols.insert(" ") |
| 13 | for s in string.ascii_letters: |
| 14 | symbols.insert(s) |
| 15 | n = symbols.num_keys() |
| 16 | merges = BPEMerges() |
| 17 | |
| 18 | tokenizer = BPETokenizer(symbols, merges) |
| 19 | |
| 20 | self.assertEqual(tokenizer.tokenize("abcd"), [1, 2, 3, 4]) |
| 21 | |
| 22 | merges.add("a", "b", n + 1) |
| 23 | self.assertEqual(tokenizer.tokenize("abcd"), [n + 1, 3, 4]) |
| 24 | |
| 25 | merges.add("c", "d", n + 2) |
| 26 | merges.add("b", "cd", n + 3) |
| 27 | self.assertEqual(tokenizer.tokenize("abcd"), [n + 1, n + 2]) |
| 28 | |
| 29 | |
| 30 | if __name__ == "__main__": |