Download scripts/02_model/train_tokenizer.py from XiaoyanLi/GooGooLM: direct link, hf CLI and curl.
- Browser
- Download file 6.76 kB
-
https://huggingface.co/XiaoyanLi/GooGooLM/resolve/main/scripts/02_model/train_tokenizer.py
- Command line
-
hf download hf://XiaoyanLi/GooGooLM/scripts/02_model/train_tokenizer.py
-
curl -L -o train_tokenizer.py https://huggingface.co/XiaoyanLi/GooGooLM/resolve/main/scripts/02_model/train_tokenizer.py
6.76 kB
| """ | |
| 训练 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" | |
| # 与官方 baseline 完全一致的特殊 token 顺序(id 固定) | |
| SPECIAL_TOKENS = ["<unk>", "<s>", "</s>", "<pad>", "<mask>"] | |
| def build_tokenizer(vocab_size: int) -> tuple[Tokenizer, BpeTrainer]: | |
| """构造与官方 baseline 相同结构的 tokenizer + trainer。""" | |
| # ── 1. Normalizer ──────────────────────────────────────────────────────── | |
| normalizer = Sequence([ | |
| Prepend(prepend=" "), | |
| NFKC(), | |
| Replace(Regex(r"\n"), "\n "), # 换行后加空格,保持词边界 | |
| Replace(Regex(r" *\n"), "\n"), # 去掉换行前多余空格 | |
| ]) | |
| # ── 2. Pre-tokenizer ───────────────────────────────────────────────────── | |
| # GPT-4 / cl100k 风格的 Unicode-aware 正则切分 | |
| 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"), # 最长24字符截断 | |
| ]) | |
| # ── 3. Tokenizer + Trainer ─────────────────────────────────────────────── | |
| 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)) | |
| # 用 transformers 包装,补充 tokenizer_config.json | |
| 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() | |