| """ |
| 训练 BPE Tokenizer |
| |
| 与官方 baseline (ltg/gpt-bert-babylm-small) 结构完全一致: |
| - Normalizer: Prepend空格 + NFKC + 换行处理 |
| - Pre-tokenizer: GPT-4风格regex切分 + ByteLevel + 最长24字符截断 |
| - Model: BPE, vocab_size=8192 |
| - 特殊token: <unk>=0, <s>=1, </s>=2, <pad>=3, <mask>=4 |
| - Post-processor: 句首自动加 <s> |
| |
| 输入: data/8_sample_B/train.txt |
| 输出: models/tokenizer/ |
| |
| 用法: |
| python scripts/02_model/train_tokenizer.py |
| python scripts/02_model/train_tokenizer.py --vocab_size 8192 |
| python scripts/02_model/train_tokenizer.py --input data/8_sample_B/train.txt |
| """ |
|
|
| import argparse |
| import json |
| from pathlib import Path |
|
|
| from tokenizers import Tokenizer, AddedToken |
| from tokenizers.models import BPE |
| from tokenizers.trainers import BpeTrainer |
| from tokenizers.normalizers import Sequence, Prepend, NFKC, Replace |
| from tokenizers.pre_tokenizers import Sequence as PreSeq, Split, ByteLevel |
| from tokenizers.processors import TemplateProcessing |
| from tokenizers import Regex |
| from transformers import PreTrainedTokenizerFast |
|
|
| ROOT = Path(__file__).parent.parent.parent |
| DEFAULT_INPUT = ROOT / "data/8_sample_B/train.txt" |
| DEFAULT_OUT = ROOT / "models/tokenizer" |
|
|
| |
| SPECIAL_TOKENS = ["<unk>", "<s>", "</s>", "<pad>", "<mask>"] |
|
|
|
|
| def build_tokenizer(vocab_size: int) -> tuple[Tokenizer, BpeTrainer]: |
| """构造与官方 baseline 相同结构的 tokenizer + trainer。""" |
|
|
| |
| normalizer = Sequence([ |
| Prepend(prepend=" "), |
| NFKC(), |
| Replace(Regex(r"\n"), "\n "), |
| Replace(Regex(r" *\n"), "\n"), |
| ]) |
|
|
| |
| |
| GPT4_REGEX = ( |
| r"[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]*" |
| r"[\p{Ll}\p{Lm}\p{Lo}\p{M}]+" |
| r"|[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]+" |
| r"[\p{Ll}\p{Lm}\p{Lo}\p{M}]*" |
| r"| ?\p{N}" |
| r"| ?[^\s\p{L}\p{N}]+[\r\n/]*" |
| r"|\s*[\r\n]+" |
| r"|\s+(?!\S)" |
| r"|\s+" |
| ) |
| pre_tokenizer = PreSeq([ |
| Split(pattern=Regex(GPT4_REGEX), behavior="isolated"), |
| ByteLevel(add_prefix_space=False, trim_offsets=True, use_regex=False), |
| Split(pattern=Regex(r".{1,24}"), behavior="isolated"), |
| ]) |
|
|
| |
| tokenizer = Tokenizer(BPE(unk_token="<unk>")) |
| tokenizer.normalizer = normalizer |
| tokenizer.pre_tokenizer = pre_tokenizer |
|
|
| trainer = BpeTrainer( |
| vocab_size=vocab_size, |
| special_tokens=SPECIAL_TOKENS, |
| min_frequency=2, |
| show_progress=True, |
| ) |
|
|
| return tokenizer, trainer |
|
|
|
|
| def add_post_processor(tokenizer: Tokenizer) -> None: |
| """添加 post-processor:句首自动插入 <s>(id=1)。""" |
| tokenizer.post_processor = TemplateProcessing( |
| single="<s> $A", |
| pair="<s> $A <s> $B", |
| special_tokens=[("<s>", tokenizer.token_to_id("<s>"))], |
| ) |
|
|
|
|
| def verify_special_token_ids(tokenizer: Tokenizer) -> None: |
| """校验特殊 token ID 与官方 baseline 一致。""" |
| expected = {"<unk>": 0, "<s>": 1, "</s>": 2, "<pad>": 3, "<mask>": 4} |
| ok = True |
| for token, expected_id in expected.items(): |
| actual_id = tokenizer.token_to_id(token) |
| status = "✅" if actual_id == expected_id else "❌" |
| print(f" {status} {token:10s} expected={expected_id} actual={actual_id}") |
| if actual_id != expected_id: |
| ok = False |
| if not ok: |
| raise ValueError("特殊 token ID 与官方 baseline 不一致!") |
|
|
|
|
| def save(tokenizer: Tokenizer, out_dir: Path, vocab_size: int) -> None: |
| """保存为 HuggingFace PreTrainedTokenizerFast 格式。""" |
| out_dir.mkdir(parents=True, exist_ok=True) |
|
|
| |
| raw_path = out_dir / "tokenizer.json" |
| tokenizer.save(str(raw_path)) |
|
|
| |
| fast_tok = PreTrainedTokenizerFast( |
| tokenizer_file=str(raw_path), |
| bos_token="<s>", |
| eos_token="</s>", |
| unk_token="<unk>", |
| sep_token="</s>", |
| pad_token="<pad>", |
| cls_token="<s>", |
| mask_token="<mask>", |
| ) |
| fast_tok.save_pretrained(str(out_dir)) |
| print(f"\n 保存到: {out_dir}") |
| print(f" 文件列表: {[f.name for f in sorted(out_dir.iterdir())]}") |
|
|
|
|
| def smoke_test(out_dir: Path) -> None: |
| """简单验证:加载后测试几个句子。""" |
| fast_tok = PreTrainedTokenizerFast.from_pretrained(str(out_dir)) |
| tests = [ |
| "The cat sat on the mat.", |
| "She gave him the book yesterday.", |
| "ran swimming swam running", |
| "Katherine can't help herself.", |
| ] |
| print("\n Smoke test:") |
| for t in tests: |
| tokens = fast_tok.tokenize(t) |
| print(f" {repr(t):45s} → {tokens}") |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--input", default=str(DEFAULT_INPUT), help="训练文件路径") |
| parser.add_argument("--output", default=str(DEFAULT_OUT), help="输出目录") |
| parser.add_argument("--vocab_size", default=8192, type=int, help="词表大小(默认8192)") |
| args = parser.parse_args() |
|
|
| input_path = Path(args.input) |
| out_dir = Path(args.output) |
|
|
| if not input_path.exists(): |
| raise FileNotFoundError(f"训练文件不存在: {input_path}") |
|
|
| print(f"训练 BPE Tokenizer") |
| print(f" 输入 : {input_path} ({input_path.stat().st_size / 1e6:.1f} MB)") |
| print(f" 输出 : {out_dir}") |
| print(f" vocab_size = {args.vocab_size}") |
| print(f" 特殊 token: {SPECIAL_TOKENS}") |
| print() |
|
|
| tokenizer, trainer = build_tokenizer(args.vocab_size) |
|
|
| print("训练中...") |
| tokenizer.train(files=[str(input_path)], trainer=trainer) |
| print(f"训练完成,实际 vocab size = {tokenizer.get_vocab_size()}") |
|
|
| add_post_processor(tokenizer) |
|
|
| print("\n特殊 token ID 校验:") |
| verify_special_token_ids(tokenizer) |
|
|
| save(tokenizer, out_dir, args.vocab_size) |
| smoke_test(out_dir) |
|
|
| print("\n完成!") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|