codekingpro/portable-devtools
114k
1"""This is an educational implementation of the byte pair encoding algorithm."""
2
3from __future__ import annotations
4
5import collections
6
7import regex
8
9import tiktoken
10
11
12class SimpleBytePairEncoding:
13 def __init__(self, *, pat_str: str, mergeable_ranks: dict[bytes, int]) -> None:
14 """Creates an Encoding object."""
15 # A regex pattern string that is used to split the input text
16 self.pat_str = pat_str
17 # A dictionary mapping token bytes to their ranks. The ranks correspond to merge priority
18 self.mergeable_ranks = mergeable_ranks
19
20 self._decoder = {token: token_bytes for token_bytes, token in mergeable_ranks.items()}
21 self._pat = regex.compile(pat_str)
22
23 def encode(self, text: str, visualise: str | None = "colour") -> list[int]:
24 """Encodes a string into tokens.
25
26 >>> enc.encode("hello world")
27 [388, 372]
28 """
29 # Use the regex to split the text into (approximately) words
30 words = self._pat.findall(text)
31 tokens = []
32 for word in words:
33 # Turn each word into tokens, using the byte pair encoding algorithm
34 word_bytes = word.encode("utf-8")
35 word_tokens = bpe_encode(self.mergeable_ranks, word_bytes, visualise=visualise)
36 tokens.extend(word_tokens)
37 return tokens
38
39 def decode_bytes(self, tokens: list[int]) -> bytes:
40 """Decodes a list of tokens into bytes.
41
42 >>> enc.decode_bytes([388, 372])
43 b'hello world'
44 """
45 return b"".join(self._decoder[token] for token in tokens)
46
47 def decode(self, tokens: list[int]) -> str:
48 """Decodes a list of tokens into a string.
49
50 Decoded bytes are not guaranteed to be valid UTF-8. In that case, we replace
51 the invalid bytes with the replacement character "�".
52
53 >>> enc.decode([388, 372])
54 'hello world'
55 """
56 return self.decode_bytes(tokens).decode("utf-8", errors="replace")
57
58 def decode_tokens_bytes(self, tokens: list[int]) -> list[bytes]:
59 """Decodes a list of tokens into a list of bytes.
60
61 Useful for visualising how a string is tokenised.
62
63 >>> enc.decode_tokens_bytes([388, 372])
64 [b'hello', b' world']
65 """
66 return [self._decoder[token] for token in tokens]
67
68 @staticmethod
69 def train(training_data: str, vocab_size: int, pat_str: str):
70 """Train a BPE tokeniser on some data!"""
71 mergeable_ranks = bpe_train(data=training_data, vocab_size=vocab_size, pat_str=pat_str)
72 return SimpleBytePairEncoding(pat_str=pat_str, mergeable_ranks=mergeable_ranks)
73
74 @staticmethod
75 def from_tiktoken(encoding):
76 if isinstance(encoding, str):
77 encoding = tiktoken.get_encoding(encoding)
78 return SimpleBytePairEncoding(
79 pat_str=encoding._pat_str, mergeable_ranks=encoding._mergeable_ranks
80 )
81
82
83def bpe_encode(
84 mergeable_ranks: dict[bytes, int], input: bytes, visualise: str | None = "colour"
85) -> list[int]:
86 parts = [bytes([b]) for b in input]
87 while True:
88 # See the intermediate merges play out!
89 if visualise:
90 if visualise in ["colour", "color"]:
91 visualise_tokens(parts)
92 elif visualise == "simple":
93 print(parts)
94
95 # Iterate over all pairs and find the pair we want to merge the most
96 min_idx = None
97 min_rank = None
98 for i, pair in enumerate(zip(parts[:-1], parts[1:])):
99 rank = mergeable_ranks.get(pair[0] + pair[1])
100 if rank is not None and (min_rank is None or rank < min_rank):
101 min_idx = i
102 min_rank = rank
103
104 # If there were no pairs we could merge, we're done!
105 if min_rank is None:
106 break
107 assert min_idx is not None
108
109 # Otherwise, merge that pair and leave the rest unchanged. Then repeat.
110 parts = parts[:min_idx] + [parts[min_idx] + parts[min_idx + 1]] + parts[min_idx + 2 :]
111
112 if visualise:
113 print()
114
115 tokens = [mergeable_ranks[part] for part in parts]
116 return tokens
117
118
119def bpe_train(
120 data: str, vocab_size: int, pat_str: str, visualise: str | None = "colour"
121) -> dict[bytes, int]:
122 # First, add tokens for each individual byte value
123 if vocab_size < 2**8:
124 raise ValueError("vocab_size must be at least 256, so we can encode all bytes")
125 ranks = {}
126 for i in range(2**8):
127 ranks[bytes([i])] = i
128
129 # Splinter up our data into lists of bytes
130 # data = "Hello world"
131 # words = [
132 # [b'H', b'e', b'l', b'l', b'o'],
133 # [b' ', b'w', b'o', b'r', b'l', b'd']
134 # ]
135 words: list[list[bytes]] = [
136 [bytes([b]) for b in word.encode("utf-8")] for word in regex.findall(pat_str, data)
137 ]
138
139 # Now, use our data to figure out which merges we should make
140 while len(ranks) < vocab_size:
141 # Find the most common pair. This will become our next token
142 stats = collections.Counter()
143 for piece in words:
144 for pair in zip(piece[:-1], piece[1:]):
145 stats[pair] += 1
146
147 most_common_pair = max(stats, key=lambda x: stats[x])
148 token_bytes = most_common_pair[0] + most_common_pair[1]
149 token = len(ranks)
150 # Add the new token!
151 ranks[token_bytes] = token
152
153 # Now merge that most common pair in all the words. That is, update our training data
154 # to reflect our decision to make that pair into a new token.
155 new_words = []
156 for word in words:
157 new_word = []
158 i = 0
159 while i < len(word) - 1:
160 if (word[i], word[i + 1]) == most_common_pair:
161 # We found our pair! Merge it
162 new_word.append(token_bytes)
163 i += 2
164 else:
165 new_word.append(word[i])
166 i += 1
167 if i == len(word) - 1:
168 new_word.append(word[i])
169 new_words.append(new_word)
170 words = new_words
171
172 # See the intermediate merges play out!
173 if visualise:
174 print(f"The current most common pair is {most_common_pair[0]} + {most_common_pair[1]}")
175 print(f"So we made {token_bytes} our {len(ranks)}th token")
176 if visualise in ["colour", "color"]:
177 print("Now the first fifty words in our training data look like:")
178 visualise_tokens([token for word in words[:50] for token in word])
179 elif visualise == "simple":
180 print("Now the first twenty words in our training data look like:")
181 for word in words[:20]:
182 print(word)
183 print("\n")
184
185 return ranks
186
187
188def visualise_tokens(token_values: list[bytes]) -> None:
189 background = [f"\u001b[48;5;{i}m" for i in [167, 179, 185, 77, 80, 68, 134]]
190 # If token boundaries do not occur at unicode character boundaries, it's unclear how best to
191 # visualise the token. Here, we'll just use the unicode replacement character to represent some
192 # fraction of a character.
193 unicode_token_values = [x.decode("utf-8", errors="replace") for x in token_values]
194
195 running_length = 0
196 last_color = None
197 for token in unicode_token_values:
198 color = background[running_length % len(background)]
199 if color == last_color:
200 color = background[(running_length + 1) % len(background)]
201 assert color != last_color
202 last_color = color
203 running_length += len(token)
204 print(color + token, end="")
205 print("\u001b[0m")
206
207
208def train_simple_encoding():
209 gpt2_pattern = (
210 r"""'s|'t|'re|'ve|'m|'ll|'d| ?[\p{L}]+| ?[\p{N}]+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""
211 )
212 with open(__file__) as f:
213 data = f.read()
214
215 enc = SimpleBytePairEncoding.train(data, vocab_size=600, pat_str=gpt2_pattern)
216
217 print("This is the sequence of merges performed in order to encode 'hello world':")
218 tokens = enc.encode("hello world")
219 assert enc.decode(tokens) == "hello world"
220 assert enc.decode_bytes(tokens) == b"hello world"
221 assert enc.decode_tokens_bytes(tokens) == [b"hello", b" world"]
222
223 return enc
224 