#!/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()