Alhdrawi commited on
Commit
ad90b77
·
verified ·
1 Parent(s): 47c37af

Upload simple_tokenizer.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. simple_tokenizer.py +156 -0
simple_tokenizer.py ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ MIT License
3
+
4
+ Copyright (c) 2021 OpenAI
5
+
6
+ Permission is hereby granted, free of charge, to any person obtaining a copy
7
+ of this software and associated documentation files (the "Software"), to deal
8
+ in the Software without restriction, including without limitation the rights
9
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
10
+ copies of the Software, and to permit persons to whom the Software is
11
+ furnished to do so, subject to the following conditions:
12
+
13
+ The above copyright notice and this permission notice shall be included in all
14
+ copies or substantial portions of the Software.
15
+
16
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
17
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
18
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
19
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
20
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
21
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
22
+ SOFTWARE.
23
+
24
+ """
25
+ import gzip
26
+ import html
27
+ import os
28
+ from functools import lru_cache
29
+
30
+ import ftfy
31
+ import regex as re
32
+
33
+
34
+ @lru_cache()
35
+ def default_bpe():
36
+ return os.path.join(os.path.dirname(os.path.abspath(__file__)), "bpe_simple_vocab_16e6.txt.gz")
37
+
38
+
39
+ @lru_cache()
40
+ def bytes_to_unicode():
41
+ """
42
+ Returns list of utf-8 byte and a corresponding list of unicode strings.
43
+ The reversible bpe codes work on unicode strings.
44
+ This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
45
+ When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
46
+ This is a signficant percentage of your normal, say, 32K bpe vocab.
47
+ To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
48
+ And avoids mapping to whitespace/control characters the bpe code barfs on.
49
+ """
50
+ bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
51
+ cs = bs[:]
52
+ n = 0
53
+ for b in range(2**8):
54
+ if b not in bs:
55
+ bs.append(b)
56
+ cs.append(2**8+n)
57
+ n += 1
58
+ cs = [chr(n) for n in cs]
59
+ return dict(zip(bs, cs))
60
+
61
+
62
+ def get_pairs(word):
63
+ """Return set of symbol pairs in a word.
64
+ Word is represented as tuple of symbols (symbols being variable-length strings).
65
+ """
66
+ pairs = set()
67
+ prev_char = word[0]
68
+ for char in word[1:]:
69
+ pairs.add((prev_char, char))
70
+ prev_char = char
71
+ return pairs
72
+
73
+
74
+ def basic_clean(text):
75
+ text = ftfy.fix_text(text)
76
+ text = html.unescape(html.unescape(text))
77
+ return text.strip()
78
+
79
+
80
+ def whitespace_clean(text):
81
+ text = re.sub(r'\s+', ' ', text)
82
+ text = text.strip()
83
+ return text
84
+
85
+
86
+ class SimpleTokenizer(object):
87
+ def __init__(self, bpe_path: str = default_bpe()):
88
+ self.byte_encoder = bytes_to_unicode()
89
+ self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
90
+ merges = gzip.open(bpe_path).read().decode("utf-8").split('\n')
91
+ merges = merges[1:49152-256-2+1]
92
+ merges = [tuple(merge.split()) for merge in merges]
93
+ vocab = list(bytes_to_unicode().values())
94
+ vocab = vocab + [v+'</w>' for v in vocab]
95
+ for merge in merges:
96
+ vocab.append(''.join(merge))
97
+ vocab.extend(['<|startoftext|>', '<|endoftext|>'])
98
+ self.encoder = dict(zip(vocab, range(len(vocab))))
99
+ self.decoder = {v: k for k, v in self.encoder.items()}
100
+ self.bpe_ranks = dict(zip(merges, range(len(merges))))
101
+ self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
102
+ 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)
103
+
104
+ def bpe(self, token):
105
+ if token in self.cache:
106
+ return self.cache[token]
107
+ word = tuple(token[:-1]) + ( token[-1] + '</w>',)
108
+ pairs = get_pairs(word)
109
+
110
+ if not pairs:
111
+ return token+'</w>'
112
+
113
+ while True:
114
+ bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
115
+ if bigram not in self.bpe_ranks:
116
+ break
117
+ first, second = bigram
118
+ new_word = []
119
+ i = 0
120
+ while i < len(word):
121
+ try:
122
+ j = word.index(first, i)
123
+ new_word.extend(word[i:j])
124
+ i = j
125
+ except:
126
+ new_word.extend(word[i:])
127
+ break
128
+
129
+ if word[i] == first and i < len(word)-1 and word[i+1] == second:
130
+ new_word.append(first+second)
131
+ i += 2
132
+ else:
133
+ new_word.append(word[i])
134
+ i += 1
135
+ new_word = tuple(new_word)
136
+ word = new_word
137
+ if len(word) == 1:
138
+ break
139
+ else:
140
+ pairs = get_pairs(word)
141
+ word = ' '.join(word)
142
+ self.cache[token] = word
143
+ return word
144
+
145
+ def encode(self, text):
146
+ bpe_tokens = []
147
+ text = whitespace_clean(basic_clean(text)).lower()
148
+ for token in re.findall(self.pat, text):
149
+ token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
150
+ bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
151
+ return bpe_tokens
152
+
153
+ def decode(self, tokens):
154
+ text = ''.join([self.decoder[token] for token in tokens])
155
+ text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors="replace").replace('</w>', ' ')
156
+ return text