Skip to content

Commit a71bc2a

Browse files
authored
Merge branch 'master' into stopwatch_sort
2 parents 49ec553 + a6d20ad commit a71bc2a

2 files changed

Lines changed: 132 additions & 0 deletions

File tree

‎DIRECTORY.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1522,6 +1522,7 @@
15221522
* [Booths Algorithm](strings/booths_algorithm.py)
15231523
* [Boyer Moore Horspool](strings/boyer_moore_horspool.py)
15241524
* [Boyer Moore Search](strings/boyer_moore_search.py)
1525+
* [Bpe Tokenizer](strings/bpe_tokenizer.py)
15251526
* [Camel Case To Snake Case](strings/camel_case_to_snake_case.py)
15261527
* [Can String Be Rearranged As Palindrome](strings/can_string_be_rearranged_as_palindrome.py)
15271528
* [Capitalize](strings/capitalize.py)

‎strings/bpe_tokenizer.py‎

Lines changed: 131 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,131 @@
1+
"""
2+
Byte-Pair Encoding: Subword-based tokenization algorithm used by
3+
state-of-the-art language models.
4+
5+
Wikipedia: https://en.wikipedia.org/wiki/Byte_pair_encoding
6+
"""
7+
8+
import itertools
9+
from collections import OrderedDict
10+
11+
12+
def get_byte_pair_counts(ids: list[int]) -> dict:
13+
"""Count consecutive byte-pairs of an encoded string.
14+
15+
>>> ids = [73, 32, 97, 109, 32, 74, 111, 110, 83, 110, 111, 119, 46]
16+
>>> get_byte_pair_counts(ids)
17+
{(73, 32): 1, (32, 97): 1, (97, 109): 1, (109, 32): 1, (32, 74): 1, (74, 111): 1, (111, 110): 1, (110, 83): 1, (83, 110): 1, (110, 111): 1, (111, 119): 1, (119, 46): 1}
18+
>>> ids = [2, 3, 6, 2, 3, 6, 2, 5]
19+
>>> get_byte_pair_counts(ids)
20+
{(2, 3): 2, (3, 6): 2, (6, 2): 2, (2, 5): 1}
21+
""" # noqa: E501
22+
counts: dict = {}
23+
for pair in itertools.pairwise(ids):
24+
counts[pair] = counts.get(pair, 0) + 1
25+
return counts
26+
27+
28+
def merge(ids: list[int], pair: tuple, idx: int) -> list[int]:
29+
"""Replace most occurring byte pair with new byte that is not used
30+
in the data. For utf-8 encoding, we start with 256 as the new byte
31+
32+
>>> ids = [2, 3, 6, 2, 3, 6, 2, 5]
33+
>>> pair = (2, 3)
34+
>>> idx = 256
35+
>>> merge(ids, pair, idx)
36+
[256, 6, 256, 6, 2, 5]
37+
"""
38+
new_ids = []
39+
i = 0
40+
while i < len(ids):
41+
if i < len(ids) - 1 and (ids[i] == pair[0] and ids[i + 1] == pair[1]):
42+
new_ids.append(idx)
43+
i += 2
44+
else:
45+
new_ids.append(ids[i])
46+
i += 1
47+
return new_ids
48+
49+
50+
class Tokenizer:
51+
"""Tokenize a string using the byte-pair encoding algorithm"""
52+
53+
def __init__(self, num_merges: int = 20, verbose: bool = False) -> None:
54+
self.num_merges = num_merges
55+
self.merges: dict = {}
56+
self.verbose = verbose
57+
58+
def encode(self, text: str) -> list[int]:
59+
"""Convert a string to tokens (bytes)
60+
61+
>>> t = Tokenizer()
62+
>>> text = "I am JonSnow."
63+
>>> t.encode(text)
64+
[73, 32, 97, 109, 32, 74, 111, 110, 83, 110, 111, 119, 46]
65+
66+
>>> t = Tokenizer()
67+
>>> text = ""
68+
>>> t.encode(text)
69+
[]
70+
"""
71+
text_b = text.encode("utf-8") # raw bytes
72+
tokens = list(map(int, text_b)) # convert to list of integers
73+
74+
if self.verbose:
75+
print(f"Input text: {text}")
76+
print(f"Tokens: {tokens}")
77+
78+
ids = list(tokens) # create a copy of tokens
79+
self.merges = OrderedDict() # store a mapping of merges (int, int) -> int
80+
max_merges = len(tokens) - 1
81+
num_merges = min(self.num_merges, max_merges)
82+
# start merging most frequently occurring byte pairs
83+
for i in range(num_merges):
84+
counts = get_byte_pair_counts(ids)
85+
pair = max(counts, key=counts.__getitem__)
86+
87+
if counts[pair] == 1:
88+
continue
89+
90+
idx = 256 + i # create new token for every merge step
91+
if self.verbose:
92+
print(f"Merging {pair} into a new token {idx}")
93+
ids = merge(ids, pair, idx)
94+
self.merges[pair] = idx
95+
96+
return ids
97+
98+
def decode(self, ids: list[int]) -> str:
99+
"""Convert a list of tokens to the original string
100+
101+
>>> t = Tokenizer()
102+
>>> ids = [73, 32, 97, 109, 32, 74, 111, 110, 83, 110, 111, 119, 46]
103+
>>> t.decode(ids)
104+
'I am JonSnow.'
105+
106+
>>> t = Tokenizer()
107+
>>> ids = []
108+
>>> t.decode(ids)
109+
''
110+
"""
111+
vocab = {idx: bytes([idx]) for idx in range(256)} # original vocabulary
112+
# The iteration of items should be in the order of
113+
# their insertion. This is the default behavior in Python 3
114+
# but we use an OrderedDict explicitly here
115+
for (p0, p1), idx in self.merges.items():
116+
vocab[idx] = vocab[p0] + vocab[p1]
117+
118+
if self.verbose:
119+
print("Vocabulary (after merging): {vocab}")
120+
121+
tokens = b"".join(vocab[idx] for idx in ids)
122+
# handle UnicodeDecodeError by replacing the invalid
123+
# start byte to conform to utf-8 format
124+
text = tokens.decode("utf-8", errors="replace")
125+
return text
126+
127+
128+
if __name__ == "__main__":
129+
import doctest
130+
131+
doctest.testmod()

0 commit comments

Comments
 (0)