Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
load.py173 linesDownload Raw Back to tiktoken
1from __future__ import annotations
2
3import base64
4import hashlib
5import os
6
7
8def read_file(blobpath: str) -> bytes:
9    if "://" not in blobpath:
10        with open(blobpath, "rb", buffering=0) as f:
11            return f.read()
12
13    if blobpath.startswith(("http://", "https://")):
14        # avoiding blobfile for public files helps avoid auth issues, like MFA prompts.
15        import requests
16
17        resp = requests.get(blobpath)
18        resp.raise_for_status()
19        return resp.content
20
21    try:
22        import blobfile
23    except ImportError as e:
24        raise ImportError(
25            "blobfile is not installed. Please install it by running `pip install blobfile`."
26        ) from e
27    with blobfile.BlobFile(blobpath, "rb") as f:
28        return f.read()
29
30
31def check_hash(data: bytes, expected_hash: str) -> bool:
32    actual_hash = hashlib.sha256(data).hexdigest()
33    return actual_hash == expected_hash
34
35
36def read_file_cached(blobpath: str, expected_hash: str | None = None) -> bytes:
37    user_specified_cache = True
38    if "TIKTOKEN_CACHE_DIR" in os.environ:
39        cache_dir = os.environ["TIKTOKEN_CACHE_DIR"]
40    elif "DATA_GYM_CACHE_DIR" in os.environ:
41        cache_dir = os.environ["DATA_GYM_CACHE_DIR"]
42    else:
43        import tempfile
44
45        cache_dir = os.path.join(tempfile.gettempdir(), "data-gym-cache")
46        user_specified_cache = False
47
48    if cache_dir == "":
49        # disable caching
50        return read_file(blobpath)
51
52    cache_key = hashlib.sha1(blobpath.encode()).hexdigest()
53
54    cache_path = os.path.join(cache_dir, cache_key)
55    if os.path.exists(cache_path):
56        with open(cache_path, "rb", buffering=0) as f:
57            data = f.read()
58        if expected_hash is None or check_hash(data, expected_hash):
59            return data
60
61        # the cached file does not match the hash, remove it and re-fetch
62        try:
63            os.remove(cache_path)
64        except OSError:
65            pass
66
67    contents = read_file(blobpath)
68    if expected_hash and not check_hash(contents, expected_hash):
69        raise ValueError(
70            f"Hash mismatch for data downloaded from {blobpath} (expected {expected_hash}). "
71            f"This may indicate a corrupted download. Please try again."
72        )
73
74    import uuid
75
76    try:
77        os.makedirs(cache_dir, exist_ok=True)
78        tmp_filename = cache_path + "." + str(uuid.uuid4()) + ".tmp"
79        with open(tmp_filename, "wb") as f:
80            f.write(contents)
81        os.rename(tmp_filename, cache_path)
82    except OSError:
83        # don't raise if we can't write to the default cache, e.g. issue #75
84        if user_specified_cache:
85            raise
86
87    return contents
88
89
90def data_gym_to_mergeable_bpe_ranks(
91    vocab_bpe_file: str,
92    encoder_json_file: str,
93    vocab_bpe_hash: str | None = None,
94    encoder_json_hash: str | None = None,
95    clobber_one_byte_tokens: bool = False,
96) -> dict[bytes, int]:
97    # NB: do not add caching to this function
98    rank_to_intbyte = [b for b in range(2**8) if chr(b).isprintable() and chr(b) != " "]
99
100    data_gym_byte_to_byte = {chr(b): b for b in rank_to_intbyte}
101    n = 0
102    for b in range(2**8):
103        if b not in rank_to_intbyte:
104            rank_to_intbyte.append(b)
105            data_gym_byte_to_byte[chr(2**8 + n)] = b
106            n += 1
107    assert len(rank_to_intbyte) == 2**8
108
109    # vocab_bpe contains the merges along with associated ranks
110    vocab_bpe_contents = read_file_cached(vocab_bpe_file, vocab_bpe_hash).decode()
111    bpe_merges = [tuple(merge_str.split()) for merge_str in vocab_bpe_contents.split("\n")[1:-1]]
112
113    def decode_data_gym(value: str) -> bytes:
114        return bytes(data_gym_byte_to_byte[b] for b in value)
115
116    # add the single byte tokens
117    # if clobber_one_byte_tokens is True, we'll replace these with ones from the encoder json
118    bpe_ranks = {bytes([b]): i for i, b in enumerate(rank_to_intbyte)}
119    del rank_to_intbyte
120
121    # add the merged tokens
122    n = len(bpe_ranks)
123    for first, second in bpe_merges:
124        bpe_ranks[decode_data_gym(first) + decode_data_gym(second)] = n
125        n += 1
126
127    import json
128
129    # check that the encoder file matches the merges file
130    # this sanity check is important since tiktoken assumes that ranks are ordered the same
131    # as merge priority
132    encoder_json = json.loads(read_file_cached(encoder_json_file, encoder_json_hash))
133    encoder_json_loaded = {decode_data_gym(k): v for k, v in encoder_json.items()}
134    # drop these two special tokens if present, since they're not mergeable bpe tokens
135    encoder_json_loaded.pop(b"<|endoftext|>", None)
136    encoder_json_loaded.pop(b"<|startoftext|>", None)
137
138    if clobber_one_byte_tokens:
139        for k in encoder_json_loaded:
140            if len(k) == 1:
141                bpe_ranks[k] = encoder_json_loaded[k]
142
143    assert bpe_ranks == encoder_json_loaded
144
145    return bpe_ranks
146
147
148def dump_tiktoken_bpe(bpe_ranks: dict[bytes, int], tiktoken_bpe_file: str) -> None:
149    try:
150        import blobfile
151    except ImportError as e:
152        raise ImportError(
153            "blobfile is not installed. Please install it by running `pip install blobfile`."
154        ) from e
155    with blobfile.BlobFile(tiktoken_bpe_file, "wb") as f:
156        for token, rank in sorted(bpe_ranks.items(), key=lambda x: x[1]):
157            f.write(base64.b64encode(token) + b" " + str(rank).encode() + b"\n")
158
159
160def load_tiktoken_bpe(tiktoken_bpe_file: str, expected_hash: str | None = None) -> dict[bytes, int]:
161    # NB: do not add caching to this function
162    contents = read_file_cached(tiktoken_bpe_file, expected_hash)
163    ret = {}
164    for line in contents.splitlines():
165        if not line:
166            continue
167        try:
168            token, rank = line.split()
169            ret[base64.b64decode(token)] = int(rank)
170        except Exception as e:
171            raise ValueError(f"Error parsing line {line!r} in {tiktoken_bpe_file}") from e
172    return ret
173 
codekingpro/portable-devtools · Team Ai