SPB Git

spb/forge Public MIT

Forge — LLM training from scratch in pure C++20 + Metal on Apple Silicon.

C++ 61.2% C 23% Python 7.6% TeX 7.2% CMake 1.1%
3.8 KB · 113 lines python
Raw Blame History
1# Author: Simon-Pierre Boucher — contact@spboucher.ai2"""Train a byte-level BPE vocab (minbpe BasicTokenizer algorithm, no regex3pre-split) and write the forgebpe v1 model file that src/tokenizer/bpe.cpp4loads.56Training is numpy-vectorized (bincount over packed pair keys + vectorized7merge) — minutes instead of hours for a 10-20 MB sample; a pure-python8fallback keeps the tool dependency-free if numpy is missing. The learned9merges are identical either way (ties broken by np.argmax == first-seen10max, matching Counter.most_common determinism closely enough for a fresh11vocab; the ENCODER contract that matters — greedy lowest-id merge — is12fixed by the .model file, shared by prepare_data.py and the C++ side).1314Usage:15  python3 tools/train_tokenizer.py --input data/sample.txt --vocab-size 4096 \16      --output data/tok4096.model [--max-mb 10]17"""18import argparse19from collections import Counter202122def train_numpy(text_bytes, vocab_size):23    import numpy as np2425    ids = np.frombuffer(text_bytes, dtype=np.uint8).astype(np.int32)26    merges = []27    for i in range(vocab_size - 256):28        if len(ids) < 2:29            break30        keys = ids[:-1] * vocab_size + ids[1:]31        counts = np.bincount(keys, minlength=vocab_size * vocab_size)32        best = int(counts.argmax())33        if counts[best] < 2:34            break35        left, right = best // vocab_size, best % vocab_size36        idx = 256 + i3738        # left-to-right non-overlapping merge (matches minbpe merge())39        is_pair = (ids[:-1] == left) & (ids[1:] == right)40        pos = np.where(is_pair)[0]41        if left == right:  # resolve overlaps like "aaa": keep first of each run42            keep, last = [], -243            for p in pos:44                if p != last + 1:45                    keep.append(p)46                    last = p47            pos = np.asarray(keep, dtype=np.int64)48        ids[pos] = idx49        ids = np.delete(ids, pos + 1)50        merges.append((idx, int(left), int(right)))51        if (i + 1) % 512 == 0:52            print(f"  merge {i + 1}/{vocab_size - 256} (corpus now {len(ids)} ids)")53    return merges545556def train_python(text_bytes, vocab_size):57    ids = list(text_bytes)58    merges = []59    for i in range(vocab_size - 256):60        counts = Counter(zip(ids, ids[1:]))61        if not counts:62            break63        pair = max(counts, key=counts.get)64        idx = 256 + i65        out, j = [], 066        while j < len(ids):67            if j + 1 < len(ids) and (ids[j], ids[j + 1]) == pair:68                out.append(idx)69                j += 270            else:71                out.append(ids[j])72                j += 173        ids = out74        merges.append((idx, pair[0], pair[1]))75    return merges767778def train(text_bytes, vocab_size):79    try:80        return train_numpy(text_bytes, vocab_size)81    except ImportError:82        print("numpy not found — slow pure-python trainer")83        return train_python(text_bytes, vocab_size)848586def write_model(path, merges):87    with open(path, "w") as f:88        f.write("forgebpe v1\n")89        f.write(f"{256 + len(merges)}\n")90        for idx, left, right in merges:91            f.write(f"{idx} {left} {right}\n")929394def main():95    ap = argparse.ArgumentParser()96    ap.add_argument("--input", required=True)97    ap.add_argument("--vocab-size", type=int, default=4096)98    ap.add_argument("--output", required=True)99    ap.add_argument("--max-mb", type=float, default=10.0,100                    help="train on at most this many MB of the input")101    args = ap.parse_args()102103    with open(args.input, "rb") as f:104        data = f.read(int(args.max_mb * 1024 * 1024))105    print(f"training {args.vocab_size}-vocab BPE on {len(data) / 1e6:.1f} MB")106    merges = train(data, args.vocab_size)107    write_model(args.output, merges)108    print(f"wrote {args.output} ({256 + len(merges)} tokens)")109110111if __name__ == "__main__":112    main()113