File size: 6,761 Bytes
83112d8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
"""
训练 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()