File size: 4,897 Bytes
9d2b68b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
from pathlib import Path
import regex as re
import json
from collections import defaultdict
from collections.abc import Iterable, Iterator


def gpt2_bytes_to_unicode() -> dict[int, str]:
        bs = (list(range(ord("!"), ord("~") + 1))
            + list(range(ord("¡"), ord("¬") + 1))
            + list(range(ord("®"), ord("ÿ") + 1)))
        cs = bs[:]
        n = 0
        for i in range(256):
            if i not in bs:
                bs.append(i)
                cs.append(256+n)
                n += 1

        return {b:chr(c) for b,c in zip(bs,cs)}
    
def gpt2_unicode_to_bytes() -> dict[str,int]:
        return {v:k for k,v in gpt2_bytes_to_unicode().items()}


PAT = r"""'(?:[sdmt]|ll|ve|re)| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""

class Tokenizer:
    def __init__(self, vocab: dict[int, bytes], merges: list[tuple[bytes,bytes]], special_tokens: list[str]| None = None):
        self.vocab = vocab
        self.merges = merges
        self.special_tokens = special_tokens or []
        self.byte_to_id = {v: k for k, v in self.vocab.items()} # for encoding
        self.merge_rank = {pair: i for i, pair in enumerate(merges)}


    
    # Load vocab from file. 
    @classmethod
    def from_file(cls, vocab_filepath: str | Path, merges_filepath: str| Path, special_tokens: list[str] | None = None) -> "Tokenizer":
        # populate vocab
        byte_decoder = gpt2_unicode_to_bytes()
        to_bytes = lambda s: bytes(byte_decoder[ch] for ch in s)

        with open(vocab_filepath, encoding = "utf-8") as f:
            raw = json.load(f)
            vocab = {k:to_bytes(v) for v,k in raw.items()}

        # Populate merges
        merges = []
        with open(merges_filepath, encoding = "utf-8") as f:
            for line in f:
                line = line.rstrip("\n")
                a , b = line.split(" ")
                merges.append((to_bytes(a), to_bytes(b)))

        return cls(vocab, merges, special_tokens)
        

    def encode(self, string: str) -> list[int]:
        # Segment on special tokens first.
        if self.special_tokens:
            specials = sorted(self.special_tokens, key = len, reverse=True)
            segment_pattern = re.compile("(" + "|".join(re.escape(t) for t in specials) + ")")
            segments = re.split(segment_pattern, string)
        else:
            segments = [string]

        result = []
        chunk_pattern = re.compile(PAT)

        for segment in segments:
            if segment in self.special_tokens:
                result.append(self.byte_to_id[segment.encode("utf-8")])

            else:
                
                for match in chunk_pattern.finditer(segment):
                    token = match.group()
                    byte_list = [bytes([x]) for x in token.encode("utf-8")]
                    new_list = byte_list.copy()

                    # Apply merges
                    while True:
                        # If we don't have enough for a pair, we break off
                        if len(new_list) < 2:
                            break

                        # Get all pairs
                        pairs = [(i1,i2) for i1,i2 in zip(new_list, new_list[1:])]

                        # Get all ranks in the merge list. 
                        rank = {p: self.merge_rank[p] for p in pairs if p in self.merge_rank}
                        if not rank:
                            break
                        min_pair = min(rank, key=rank.get)

                        # We apply merge now
                        merge_applied_list = []
                        i = 0
                        while i < len(new_list):
                            if i < len(new_list) - 1 and new_list[i] == min_pair[0] and new_list[i+1] == min_pair[1]:
                                merge_applied_list.append(min_pair[0] + min_pair[1]) # Concatenate/merge
                                i += 2
                            else:
                                merge_applied_list.append(new_list[i])
                                i += 1

                        new_list = merge_applied_list

                    for tok in new_list:
                        result.append(self.byte_to_id[tok])


        return result 

    
    def encode_iterable(self, iterable: Iterable[str]) -> Iterator[int]:
        for chunk in iterable:
            yield from self.encode(chunk)



    def decode(self, byte_array: list[int]) -> str:
        decoded: bytes = b""
        for i in byte_array:
            decoded += self.vocab[i]
        return decoded.decode("utf-8", errors="replace")



if __name__ == "__main__":
     #tokenizer = Tokenizer.from_file( "tests/fixtures/train-bpe-reference-vocab.json","tests/fixtures/train-bpe-reference-merges.txt")
     #print(tokenizer.encode("hello my friend"))
     #print(tokenizer.decode([259, 76, 491, 486, 377, 73, 69, 269]))
     ...