Files
tei-1/train_tokenizer.py
T
2026-06-06 17:06:08 -05:00

72 lines
2.4 KiB
Python

#!/usr/bin/env python3
"""Train a custom 128k byte-level BPE tokenizer from FineWeb-EDU."""
from __future__ import annotations
import argparse
from pathlib import Path
from datasets import load_dataset
from tokenizers import Tokenizer, decoders, models, pre_tokenizers, processors, trainers
from transformers import PreTrainedTokenizerFast
def text_iterator(dataset_name: str, dataset_config: str | None, split: str, text_column: str, max_docs: int | None):
ds = load_dataset(dataset_name, dataset_config, split=split, streaming=True)
n = 0
for row in ds:
text = row.get(text_column)
if text:
yield text
n += 1
if max_docs is not None and n >= max_docs:
break
def main() -> None:
p = argparse.ArgumentParser()
p.add_argument("--out_dir", type=str, required=True)
p.add_argument("--vocab_size", type=int, default=128_000)
p.add_argument("--dataset_name", type=str, default="HuggingFaceFW/fineweb-edu")
p.add_argument("--dataset_config", type=str, default="sample-10BT")
p.add_argument("--split", type=str, default="train")
p.add_argument("--text_column", type=str, default="text")
p.add_argument("--max_docs", type=int, default=2_000_000, help="Cap docs used for tokenizer training; set 0 for no cap")
args = p.parse_args()
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
max_docs = None if args.max_docs == 0 else args.max_docs
eos = "<|endoftext|>"
unk = "<|unk|>"
tokenizer = Tokenizer(models.BPE(unk_token=unk))
tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False)
tokenizer.decoder = decoders.ByteLevel()
tokenizer.post_processor = processors.ByteLevel(trim_offsets=False)
trainer = trainers.BpeTrainer(
vocab_size=args.vocab_size,
min_frequency=2,
show_progress=True,
special_tokens=[unk, eos],
)
tokenizer.train_from_iterator(
text_iterator(args.dataset_name, args.dataset_config, args.split, args.text_column, max_docs),
trainer=trainer,
)
fast = PreTrainedTokenizerFast(
tokenizer_object=tokenizer,
unk_token=unk,
bos_token=eos,
eos_token=eos,
pad_token=eos,
)
fast.save_pretrained(out_dir)
print(f"saved tokenizer with vocab size {len(fast):,} to {out_dir}")
if __name__ == "__main__":
main()