| 32 | |
| 33 | |
| 34 | class Encoder: |
| 35 | def __init__(self, encoder, bpe_merges): |
| 36 | self.encoder = encoder |
| 37 | self.decoder = {v: k for k, v in self.encoder.items()} |
| 38 | self.bpe_ranks = dict(zip(bpe_merges, range(len(bpe_merges)))) |
| 39 | self.cache = {} |
| 40 | self.max_len = 0 |
| 41 | |
| 42 | def bpe(self, token): |
| 43 | if token in self.cache: |
| 44 | return self.cache[token] |
| 45 | word = tuple(token) |
| 46 | pairs = get_pairs(word) |
| 47 | if not pairs: |
| 48 | return token |
| 49 | |
| 50 | while True: |
| 51 | bigram = min(pairs, key=lambda pair: self.bpe_ranks.get(pair, float('inf'))) |
| 52 | if bigram not in self.bpe_ranks: |
| 53 | break |
| 54 | first, second = bigram |
| 55 | new_word = [] |
| 56 | i = 0 |
| 57 | while i < len(word): |
| 58 | try: |
| 59 | j = word.index(first, i) |
| 60 | new_word.extend(word[i:j]) |
| 61 | i = j |
| 62 | except: |
| 63 | new_word.extend(word[i:]) |
| 64 | break |
| 65 | |
| 66 | if word[i] == first and i < len(word) - 1 and word[i + 1] == second: |
| 67 | new_word.append(first + second) |
| 68 | i += 2 |
| 69 | else: |
| 70 | new_word.append(word[i]) |
| 71 | i += 1 |
| 72 | new_word = tuple(new_word) |
| 73 | word = new_word |
| 74 | if len(word) == 1: |
| 75 | break |
| 76 | else: |
| 77 | pairs = get_pairs(word) |
| 78 | word = ' '.join(word) |
| 79 | self.cache[token] = word |
| 80 | return word |
| 81 | |
| 82 | def encode(self, text): |
| 83 | return [self.encoder.get(token, 1) for token in self.tokenize(text)] |
| 84 | |
| 85 | def decode(self, tokens): |
| 86 | text = ''.join([self.decoder[token] for token in tokens]) |
| 87 | return text |
| 88 | |
| 89 | def tokenize(self, text): |
| 90 | bpe_tokens = [] |
| 91 | bpe_tokens.extend(bpe_token for bpe_token in self.bpe(text).split(' ')) |