gollem-v5-ckpts / train_tokenizer.py
Maggio33's picture
Upload train_tokenizer.py with huggingface_hub
2a8d44e verified
Raw History Blame
2.88 kB
#!/usr/bin/env python3
"""train_tokenizer.py - GoLLeM-v5 (EN) tokenizer training recipe.
Reproduces the canonical `tokenizer.json` shipped in this repo:
- byte-level BPE (GPT-2 lineage), pre-tokenizer + decoder = ByteLevel
- model vocab 12285 (256 byte alphabet + 12029 learned merges)
- 3 special tokens appended: <|endoftext|> <|im_start|> <|im_end|>
(=> 12288 effective ids; models pad the embedding to 12288)
- trained on the minimal-en-corpus (FineWeb-Edu EN broad mix; `en.parquet`,
the same source uploaded as SlayerLab/minimal-en-corpus-5b).
The shipped `tokenizer.json` remains the canonical / authoritative artifact:
a fresh run reproduces a functionally-equivalent tokenizer, but exact merge
order depends on the corpus snapshot/order, so byte-identity is not guaranteed.
Use this script to audit the method; use `tokenizer.json` for exact parity.
Usage:
python train_tokenizer.py --corpus en.parquet --out tokenizer.json
python train_tokenizer.py --limit 50000 --out tok_smoke.json # quick recipe check
"""
import argparse
import pyarrow.parquet as pq
from tokenizers import Tokenizer, models, trainers, pre_tokenizers, decoders
SPECIAL = ["<|endoftext|>", "<|im_start|>", "<|im_end|>"]
def iter_text(parquet_path, col="text", limit=None):
pf = pq.ParquetFile(parquet_path)
n = 0
for batch in pf.iter_batches(batch_size=10000, columns=[col]):
for t in batch.column(0).to_pylist():
if not t:
continue
yield t
n += 1
if limit and n >= limit:
return
def build(vocab_size):
tok = Tokenizer(models.BPE())
# ByteLevel with add_prefix_space=False matches the canonical pre_tokenizer/decoder.
tok.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
tok.decoder = decoders.ByteLevel()
trainer = trainers.BpeTrainer(
vocab_size=vocab_size,
special_tokens=SPECIAL,
initial_alphabet=pre_tokenizers.ByteLevel.alphabet(), # full 256-byte alphabet
show_progress=True,
)
return tok, trainer
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--corpus", default="C:/Projekty/datasets/build/en/en.parquet")
ap.add_argument("--out", default="tokenizer.json")
ap.add_argument("--col", default="text")
# 12288 = 12285 learned (256 bytes + 12029 merges) + 3 specials, as in canonical.
ap.add_argument("--vocab", type=int, default=12288)
ap.add_argument("--limit", type=int, default=None, help="cap #docs (smoke test)")
a = ap.parse_args()
tok, trainer = build(a.vocab)
tok.train_from_iterator(iter_text(a.corpus, a.col, a.limit), trainer=trainer)
tok.save(a.out)
print(f"saved {a.out} vocab_size={tok.get_vocab_size()} "
f"(specials={SPECIAL}, pre_tokenizer=ByteLevel)")
if __name__ == "__main__":
main()