| """txt2ids / ids2txt — the runtime interface between text and the |
| binary token stream the model consumes. |
| |
| - encode: text (or the tagged stream) -> u8 ids, longest-match, C-speed: |
| multi-byte tokens replaced by sentinel bytes, then a 256-entry translate. |
| - decode: u8 ids -> text, VERBATIM (tags stay literal strings). |
| - display: verbatim text -> human view (tags become newline/space/tab, |
| <bos>/<eos>/<pad>/<msk> dropped). |
| |
| roundtrip: tokenize, decode, re-encode, assert byte-identical ids, and |
| assert the human views of source and decoded text match. |
| """ |
| import multiprocessing as mp |
| import sys |
| import time |
| from pathlib import Path |
|
|
| WORKERS = 12 |
| DATA = Path(__file__).resolve().parent.parent / "data" |
| SRC = Path(__file__).resolve().parent.parent / "src" |
|
|
| PIECES = None |
| TOKENS = None |
| TAG_ID = {} |
| CHAR_ID = {} |
| TABLE_ENC = None |
| SENT_ENC = [] |
| TABLE_DEC = None |
| SENT_DEC = [] |
|
|
|
|
| def load_vocab(): |
| global TOKENS, TAG_ID, CHAR_ID, TABLE_ENC, SENT_ENC |
| raw = (DATA / "vocab.bin").read_bytes() |
| toks, p = [], 0 |
| while p < len(raw): |
| n = raw[p] |
| toks.append(raw[p + 1:p + 1 + n]) |
| p += 1 + n |
| txt = (SRC / "vocab.txt").read_bytes().split(b"\n") |
| if txt and txt[-1] == b"": |
| txt.pop() |
| assert toks == txt, "vocab.bin does not match vocab.txt" |
| TOKENS = toks |
| for i, t in enumerate(toks): |
| if t.startswith(b"<") and t.endswith(b">"): |
| TAG_ID[t] = i |
| else: |
| CHAR_ID[t] = i |
| |
| table = bytearray([TAG_ID[b"<unk>"]] * 256) |
| multi = [(t, i) for t, i in TAG_ID.items()] + \ |
| [(t, i) for t, i in CHAR_ID.items() if len(t) > 1] |
| sent = 200 |
| for t, i in multi: |
| table[sent] = i |
| SENT_ENC.append((t, bytes([sent]))) |
| sent += 1 |
| for t, i in CHAR_ID.items(): |
| if len(t) == 1: |
| table[t[0]] = i |
| table[10] = TAG_ID[b"<nwl>"] |
| table[32] = TAG_ID[b"<spc>"] |
| table[9] = TAG_ID[b"<tab>"] |
| TABLE_ENC = bytes(table) |
| |
| |
| global TABLE_DEC, SENT_DEC |
| dt = bytearray(b"?") * 256 |
| sent = 0xF0 |
| for i, t in enumerate(toks): |
| if len(t) == 1: |
| dt[i] = t[0] |
| else: |
| dt[i] = sent |
| SENT_DEC.append((bytes([sent]), t)) |
| sent += 1 |
| TABLE_DEC = bytes(dt) |
| print(f"vocab loaded: {len(toks)} tokens") |
|
|
|
|
| def encode_bytes(b): |
| for tok, sent in SENT_ENC: |
| b = b.replace(tok, sent) |
| return b.translate(TABLE_ENC) |
|
|
|
|
| def encode_piece_slice(lo_hi): |
| lo, hi = lo_hi |
| sep = bytes([TAG_ID[b"<eos>"], TAG_ID[b"<bos>"]]) |
| return sep.join(encode_bytes(PIECES[i]) for i in range(lo, hi)) |
|
|
|
|
| def decode_slice(lo_hi): |
| lo, hi = lo_hi |
| ids = PIECES[lo:hi] |
| out = ids.translate(TABLE_DEC) |
| for sent, tok in SENT_DEC: |
| out = out.replace(sent, tok) |
| return out |
|
|
|
|
| def run_parallel(fn, n, chunk): |
| bounds = [(i, min(i + chunk, n)) for i in range(0, n, chunk)] |
| with mp.Pool(WORKERS) as pool: |
| return pool.map(fn, bounds) |
|
|
|
|
| def text_to_ids(text): |
| assert text.startswith(b"<bos>") |
| global PIECES |
| PIECES = text[5:].split(b"<eos><bos>") |
| n = len(PIECES) |
| outs = run_parallel(encode_piece_slice, n, (n + WORKERS - 1) // WORKERS) |
| eos_id, bos_id = TAG_ID[b"<eos>"], TAG_ID[b"<bos>"] |
| sep = bytes([eos_id, bos_id]) |
| return bytes([bos_id]) + sep.join(outs) |
|
|
|
|
| def ids_to_text(ids): |
| global PIECES |
| PIECES = ids |
| n = len(ids) |
| outs = run_parallel(decode_slice, n, (n + WORKERS - 1) // WORKERS) |
| return b"".join(outs) |
|
|
|
|
| def display(text): |
| for tag, sub in ((b"<nwl>", b"\n"), (b"<spc>", b" "), (b"<tab>", b"\t"), |
| (b"<bos>", b""), (b"<eos>", b""), (b"<pad>", b""), |
| (b"<msk>", b"")): |
| text = text.replace(tag, sub) |
| return text |
|
|
|
|
| def cmd_tokenize(): |
| load_vocab() |
| t0 = time.time() |
| ids = text_to_ids((DATA / "tinystories-cleaned.bin").read_bytes()) |
| (DATA / "tinystories-ids.bin").write_bytes(ids) |
| print(f"tokenized: {len(ids):,} ids in {time.time()-t0:.1f}s -> " |
| f"tinystories-ids.bin") |
|
|
|
|
| def cmd_detokenize(): |
| load_vocab() |
| t0 = time.time() |
| text = ids_to_text((DATA / "tinystories-ids.bin").read_bytes()) |
| print(f"detokenized: {len(text):,} bytes in {time.time()-t0:.1f}s") |
| disp = display(text) |
| (DATA / "tinystories-ids.txt").write_bytes(disp) |
| print(f"display text: {len(disp):,} bytes -> tinystories-ids.txt") |
| return text |
|
|
|
|
| def cmd_roundtrip(): |
| cmd_tokenize() |
| text = cmd_detokenize() |
| orig_ids = (DATA / "tinystories-ids.bin").read_bytes() |
| t0 = time.time() |
| ids2 = text_to_ids(text) |
| print(f"re-encoded in {time.time()-t0:.1f}s") |
| print("roundtrip ids identical:", ids2 == orig_ids) |
| src = (DATA / "tinystories-cleaned.bin").read_bytes() |
| print("human view (display) identical:", display(text) == display(src)) |
| if ids2 != orig_ids: |
| i = next((k for k, (a, b) in enumerate(zip(ids2, orig_ids)) if a != b), -1) |
| print(" first diff at id", i) |
| sys.exit(1) |
|
|
|
|
| if __name__ == "__main__": |
| cmd = sys.argv[1] if len(sys.argv) > 1 else "roundtrip" |
| {"tokenize": cmd_tokenize, "detokenize": cmd_detokenize, |
| "roundtrip": cmd_roundtrip}[cmd]() |
|
|