| 75 | self.eot_token = self.encoder.get("<|endoftext|>", 49407) |
| 76 | |
| 77 | def bpe(self, token): |
| 78 | if token in self.cache: |
| 79 | return self.cache[token] |
| 80 | |
| 81 | word = tuple(token[:-1]) + (token[-1] + "</w>",) |
| 82 | pairs = get_pairs(word) |
| 83 | |
| 84 | if not pairs: |
| 85 | return token + "</w>" |
| 86 | |
| 87 | while True: |
| 88 | bigram = min(pairs, key=lambda p: self.bpe_ranks.get(p, float("inf"))) |
| 89 | if bigram not in self.bpe_ranks: |
| 90 | break |
| 91 | |
| 92 | first, second = bigram |
| 93 | new_word = [] |
| 94 | i = 0 |
| 95 | while i < len(word): |
| 96 | try: |
| 97 | j = word.index(first, i) |
| 98 | new_word.extend(word[i:j]) |
| 99 | i = j |
| 100 | except ValueError: |
| 101 | new_word.extend(word[i:]) |
| 102 | break |
| 103 | |
| 104 | if word[i] == first and i + 1 < len(word) and word[i + 1] == second: |
| 105 | new_word.append(first + second) |
| 106 | i += 2 |
| 107 | else: |
| 108 | new_word.append(word[i]) |
| 109 | i += 1 |
| 110 | |
| 111 | word = tuple(new_word) |
| 112 | if len(word) == 1: |
| 113 | break |
| 114 | pairs = get_pairs(word) |
| 115 | |
| 116 | result = " ".join(word) |
| 117 | self.cache[token] = result |
| 118 | return result |
| 119 | |
| 120 | def encode(self, text, context_length=32): |
| 121 | # Clean + lowercase |