File size: 4,665 Bytes
d8b3c96
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2cb0b71
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d8b3c96
2cb0b71
 
 
 
 
 
 
 
d8b3c96
 
 
 
2cb0b71
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d8b3c96
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Byte-level BPE compatible with babble / booper-pretrain tokenizer.json."""
from __future__ import annotations

import json
import re
from pathlib import Path

_CHUNK_RE = re.compile(r"\s+|\S+")


def _merge_ids(ids: list[int], a: int, b: int, new_id: int) -> list[int]:
    if len(ids) < 2:
        return ids
    out: list[int] = []
    i = 0
    n = len(ids)
    while i < n:
        if i + 1 < n and ids[i] == a and ids[i + 1] == b:
            out.append(new_id)
            i += 2
        else:
            out.append(ids[i])
            i += 1
    return out


class BPETokenizer:
    def __init__(self, merges: list[tuple[int, int, int]]) -> None:
        self.merges = merges
        vocab: dict[int, bytes] = {i: bytes([i]) for i in range(256)}
        for a, b, new_id in merges:
            vocab[new_id] = vocab[a] + vocab[b]
        self.vocab = vocab
        self._ranks = {(a, b): i for i, (a, b, _) in enumerate(merges)}
        self._pair_to_id = {(a, b): new_id for a, b, new_id in merges}
        base = 256 + len(merges)
        self.pad, self.bos, self.sep, self.eos = base, base + 1, base + 2, base + 3

    @property
    def vocab_size(self) -> int:
        return 256 + len(self.merges) + 4

    def _encode_chunk(self, chunk: str) -> list[int]:
        ids = list(chunk.encode("utf-8"))
        ranks = self._ranks
        pair_to_id = self._pair_to_id
        while len(ids) >= 2:
            best_rank = None
            best_pair = None
            for a, b in zip(ids, ids[1:]):
                r = ranks.get((a, b))
                if r is not None and (best_rank is None or r < best_rank):
                    best_rank = r
                    best_pair = (a, b)
            if best_pair is None:
                break
            ids = _merge_ids(ids, best_pair[0], best_pair[1], pair_to_id[best_pair])
        return ids

    def _build_fast(self):
        try:
            from tokenizers import Tokenizer
            from tokenizers import models as tokmodels
        except Exception:
            self._fast = None
            return
        id_to_tok = {i: bytes([i]).decode("latin-1") for i in range(256)}
        vocab = {s: i for i, s in id_to_tok.items()}
        hf_merges: list[tuple[str, str]] = []
        for a, b, nid in self.merges:
            sa, sb = id_to_tok[a], id_to_tok[b]
            merged = sa + sb
            id_to_tok[nid] = merged
            vocab[merged] = nid
            hf_merges.append((sa, sb))
        fast = Tokenizer(tokmodels.BPE(vocab, hf_merges, fuse_unk=False))
        self._fast = fast

    def _ensure_fast(self) -> None:
        if getattr(self, "_fast", None) is None and not hasattr(self, "_fast_tried"):
            self._fast_tried = True
            self._build_fast()

    def encode(self, text: str) -> list[int]:
        self._ensure_fast()
        if getattr(self, "_fast", None) is not None:
            ids: list[int] = []
            for chunk in _CHUNK_RE.findall(text):
                raw = chunk.encode("utf-8").decode("latin-1")
                ids.extend(self._fast.encode(raw).ids)
            return ids
        ids = []
        for chunk in _CHUNK_RE.findall(text):
            ids.extend(self._encode_chunk(chunk))
        return ids

    def encode_docs(self, texts: list[str]) -> list[int]:
        """Encode many docs and join with eos. Uses tokenizers encode_batch."""
        self._ensure_fast()
        chunks: list[str] = []
        lens: list[int] = []
        for text in texts:
            cs = _CHUNK_RE.findall(text)
            lens.append(len(cs))
            chunks.extend(c.encode("utf-8").decode("latin-1") for c in cs)
        out: list[int] = []
        if self._fast is not None and chunks:
            encs = self._fast.encode_batch(chunks)
            i = 0
            for n in lens:
                for _ in range(n):
                    out.extend(encs[i].ids)
                    i += 1
                out.append(self.eos)
            return out
        for text in texts:
            out.extend(self.encode(text))
            out.append(self.eos)
        return out

    def decode(self, ids: list[int]) -> str:
        raw = bytearray()
        for i in ids:
            piece = self.vocab.get(i)
            if piece is not None:
                raw.extend(piece)
        return bytes(raw).decode("utf-8", errors="replace")

    def to_json(self, path: Path) -> None:
        path.write_text(json.dumps({"merges": [list(m) for m in self.merges]}))

    @classmethod
    def from_json(cls, path: Path | str) -> "BPETokenizer":
        raw = json.loads(Path(path).read_text())
        return cls([tuple(m) for m in raw["merges"]])