| 60 | |
| 61 | |
| 62 | class SimpleTokenizer(object): |
| 63 | def __init__(self, bpe_path: str = default_bpe()): |
| 64 | self.byte_encoder = bytes_to_unicode() |
| 65 | self.byte_decoder = {v: k for k, v in self.byte_encoder.items()} |
| 66 | merges = gzip.open(bpe_path).read().decode("utf-8").split('\n') |
| 67 | merges = merges[1:49152-256-2+1] |
| 68 | merges = [tuple(merge.split()) for merge in merges] |
| 69 | vocab = list(bytes_to_unicode().values()) |
| 70 | vocab = vocab + [v+'</w>' for v in vocab] |
| 71 | for merge in merges: |
| 72 | vocab.append(''.join(merge)) |
| 73 | vocab.extend(['<|startoftext|>', '<|endoftext|>']) |
| 74 | self.encoder = dict(zip(vocab, range(len(vocab)))) |
| 75 | self.decoder = {v: k for k, v in self.encoder.items()} |
| 76 | self.bpe_ranks = dict(zip(merges, range(len(merges)))) |
| 77 | self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'} |
| 78 | self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", re.IGNORECASE) |
| 79 | |
| 80 | def bpe(self, token): |
| 81 | if token in self.cache: |
| 82 | return self.cache[token] |
| 83 | word = tuple(token[:-1]) + ( token[-1] + '</w>',) |
| 84 | pairs = get_pairs(word) |
| 85 | |
| 86 | if not pairs: |
| 87 | return token+'</w>' |
| 88 | |
| 89 | while True: |
| 90 | bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf'))) |
| 91 | if bigram not in self.bpe_ranks: |
| 92 | break |
| 93 | first, second = bigram |
| 94 | new_word = [] |
| 95 | i = 0 |
| 96 | while i < len(word): |
| 97 | try: |
| 98 | j = word.index(first, i) |
| 99 | new_word.extend(word[i:j]) |
| 100 | i = j |
| 101 | except: |
| 102 | new_word.extend(word[i:]) |
| 103 | break |
| 104 | |
| 105 | if word[i] == first and i < len(word)-1 and word[i+1] == second: |
| 106 | new_word.append(first+second) |
| 107 | i += 2 |
| 108 | else: |
| 109 | new_word.append(word[i]) |
| 110 | i += 1 |
| 111 | new_word = tuple(new_word) |
| 112 | word = new_word |
| 113 | if len(word) == 1: |
| 114 | break |
| 115 | else: |
| 116 | pairs = get_pairs(word) |
| 117 | word = ' '.join(word) |
| 118 | self.cache[token] = word |
| 119 | return word |
nothing calls this directly
no outgoing calls
no test coverage detected