File size: 4,861 Bytes
f4e6a29
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tokenizer de Cortex-1 para Hugging Face (equivalente exacto al BPE byte-level propio)."""
import json
import regex
from transformers import PreTrainedTokenizer

PRETOKEN = regex.compile(
    r"""'(?:[sdmt]|ll|ve|re)|[^\r\n\p{L}\p{N}]?+\p{L}+|\p{N}{1,3}| ?[^\s\p{L}\p{N}]++[\r\n]*|\s*[\r\n]|\s+(?!\S)|\s+"""
)


def bytes_to_unicode():
    """Tabla GPT-2: byte -> carácter imprimible (idéntica a la del modelo propio)."""
    bs = (list(range(ord("!"), ord("~") + 1)) + list(range(ord("¡"), ord("¬") + 1))
          + list(range(ord("®"), ord("ÿ") + 1)))
    cs = bs[:]
    n = 0
    for b in range(256):
        if b not in bs:
            bs.append(b)
            cs.append(256 + n)
            n += 1
    return dict(zip(bs, (chr(c) for c in cs)))


_B2U = bytes_to_unicode()
_U2B = {v: k for k, v in _B2U.items()}


class CortexTokenizer(PreTrainedTokenizer):
    """BPE byte-level idéntico a ``cortex.tokenizer.encoder.BPETokenizer``.

    El pre-tokenizador usa el MISMO regex (poseedores y lookahead incluidos,
    vía librería ``regex``) y los merges se aplican por orden de aprendizaje;
    el test de paridad encode(tok propio) == encode(tok HF) lo verifica.
    """

    vocab_files_names = {"vocab_file": "vocab.json", "merges_file": "merges.txt"}
    model_input_names = ["input_ids", "attention_mask"]

    def __init__(self, vocab_file=None, merges_file=None, **kwargs):
        with open(vocab_file, encoding="utf-8") as f:
            self._vocab = json.load(f)          # {token: id}
        with open(merges_file, encoding="utf-8") as f:
            merge_lines = [ln.strip() for ln in f if ln.strip() and not ln.startswith("#")]
        self._merges = {tuple(ln.split()): i for i, ln in enumerate(merge_lines)}
        self._id2tok = {i: t for t, i in self._vocab.items()}
        self._cache = {}
        super().__init__(vocab_file=vocab_file, merges_file=merges_file, **kwargs)

    @property
    def vocab_size(self):
        return len(self._vocab)

    def get_vocab(self):
        """Requerido por transformers (_add_tokens lo llama en __init__)."""
        return dict(self._vocab)

    # ------------------------------------------------------------ tokenización
    def _bpe(self, word: str):
        if word in self._cache:
            return self._cache[word]
        parts = list(word)
        while len(parts) > 1:
            pairs = list(zip(parts, parts[1:]))
            best = min(pairs, key=lambda p: self._merges.get(p, float("inf")))
            if best not in self._merges:
                break
            out, i = [], 0
            while i < len(parts):
                if i < len(parts) - 1 and (parts[i], parts[i + 1]) == best:
                    out.append(parts[i] + parts[i + 1])
                    i += 2
                else:
                    out.append(parts[i])
                    i += 1
            parts = out
        self._cache[word] = parts
        return parts

    def _split_specials(self, text: str):
        """Parte el texto conservando literales <|x|> como tokens (ids 0-8)."""
        sp = sorted((s for s in self._vocab if s.startswith("<|") and s.endswith("|>")),
                    key=len, reverse=True)
        out, i, buf = [], 0, ""
        while i < len(text):
            hit = next((s for s in sp if text.startswith(s, i)), None)
            if hit:
                if buf:
                    out.append(("text", buf)); buf = ""
                out.append(("special", hit)); i += len(hit)
            else:
                buf += text[i]; i += 1
        if buf:
            out.append(("text", buf))
        return out

    def _tokenize(self, text):
        tokens = []
        for kind, chunk in self._split_specials(text):
            if kind == "special":
                tokens.append(chunk)
                continue
            for word in PRETOKEN.findall(chunk):
                encoded = "".join(_B2U[b] for b in word.encode("utf-8"))
                tokens.extend(self._bpe(encoded))
        return tokens

    # ------------------------------------------------------------ conversión
    def _convert_token_to_id(self, token):
        return self._vocab.get(token, self._vocab.get("<|unk|>", 0))

    def _convert_id_to_token(self, index):
        return self._id2tok.get(index, "<|unk|>")

    def convert_tokens_to_string(self, tokens):
        return "".join(tokens)

    def _decode(self, token_ids, skip_special_tokens=False, **kwargs):
        pieces = []
        for i in token_ids:
            tok = self._id2tok.get(int(i))
            if tok is None:
                continue
            if skip_special_tokens and tok in self.all_special_tokens:
                continue
            pieces.append(tok)
        raw = "".join(pieces)
        data = bytes(_U2B[c] for c in raw if c in _U2B)
        return data.decode("utf-8", errors="replace")