从零预训练ZEN模型:create_pre_train_data.py数据准备全解析
【免费下载链接】ZENA BERT-based Chinese Text Encoder Enhanced by N-gram Representations项目地址: https://gitcode.com/gh_mirrors/zen10/ZEN
ZEN 是一个基于 BERT 的中文文本编码器,通过引入N-gram 表示大幅增强模型对中文词与短语的理解能力。想要从零预训练 ZEN 模型,第一步就是使用create_pre_train_data.py完成高质量的数据准备。本文将面向新手,一步步拆解 ZEN 预训练数据生成脚本的参数、流程与踩坑点,让你快速上手中文预训练。
上图是 ZEN 的整体架构:左侧为字符级编码器(Character Encoder),右侧为 N-gram 级编码器(N-gram Encoder),两者通过相加操作逐层融合,这正是 ZEN 优于纯字符模型的关键设计。数据准备脚本就是为这个双通道结构准备"字符 + N-gram"双份训练信号。
一、ZEN 预训练为什么需要专门的数据准备脚本?
与普通 BERT 不同,ZEN 的预训练目标有两个:
- 掩码语言模型(MLM):随机遮盖部分字符,让模型预测被遮盖的字;
- N-gram 增强:从词表中匹配长度为 2~7 的字符组合(如"粤港澳""大湾区"),作为额外输入注入编码器。
因此,原始语料必须先被脚本处理成包含tokens、ngram_ids、ngram_positions等字段的 JSON 样本,run_pre_train.py才能直接读取训练。这个"翻译"过程,就是create_pre_train_data.py的职责。
二、环境准备与项目获取
在运行脚本前,先确认环境满足 requirements.txt 中的依赖,核心包括 PyTorch(>=1.2)、transformers、tqdm 等。项目代码位于 ZEN 仓库根目录,先获取项目:
git clone https://gitcode.com/gh_mirrors/zen10/ZEN数据准备脚本位于 examples/create_pre_train_data.py,运行方式为:
python create_pre_train_data.py --train_corpus 语料.txt --output_dir 输出目录 --bert_model bert-base-chinese三、语料格式要求:一个空行就是一篇文章
脚本按空行划分文档,这是新手最容易忽略的一点:
- 每一行被视为一个"句子",会被分词后作为一个 segment;
- 连续的空行表示文档边界,空行之间的所有行组成一篇文档;
- 文档边界是必需的,否则无法构造"随机下一句"(Next Sentence)负样本,脚本会直接报错退出。
例如zhwiki.txt中,每篇百科条目之间留一个空行即可。相关逻辑见create_pre_train_data.py的Loading Dataset部分(examples/create_pre_train_data.py)。
四、核心参数逐一解读:最快配置方法
脚本参数不多,但每个都直接影响训练效果。下面是最常用参数的速查表:
| 参数 | 默认值 | 作用 | 建议 |
|---|---|---|---|
--train_corpus | 必填 | 原始中文语料路径 | 使用带空行分隔的多文档语料 |
--output_dir | 必填 | 预生成数据输出目录 | 与run_pre_train.py的--pregenerated_data保持一致 |
--bert_model | 必填 | 分词器来源模型 | 中文任务用bert-base-chinese |
--epochs_to_generate | 3 | 生成几轮(epoch)数据 | 语料小时可设 5~10 |
--max_seq_len | 128 | 每条样本最大长度 | 长文本任务可调至 256/512 |
--masked_lm_prob | 0.15 | 字符遮盖概率 | 与 BERT 标准一致即可 |
--max_predictions_per_seq | 20 | 每条样本最多遮盖数 | 随max_seq_len增大而增大 |
--do_whole_word_mask | 关闭 | 是否整词遮盖 | 建议开启 |
--reduce_memory | 关闭 | 用磁盘换内存 | 语料很大时务必开启 |
--max_ngram_in_sequence | 20 | 每条样本最多匹配 N-gram 数 | 默认即可 |
需要注意:脚本虽然提供了--ngram_list参数,但实际加载 N-gram 词表时使用的是ZenNgramDict(args.bert_model, ...),即从bert_model目录下的ngram.txt读取(见 ZEN/ngram_utils.py)。因此请确保ngram.txt与模型文件放在同一目录。
五、数据生成流程拆解:脚本内部做了什么?
整个处理链路可以概括为四步:
- 加载并分词:逐行读取语料,用
BertTokenizer将每句转为 token 列表,按空行切分为文档,存入DocumentDatabase; - 构造句子对:参照 BERT 的逻辑,把文档切成 A/B 两段,50% 概率从其他文档随机采样作为负样本,构造下一句预测任务;
- 生成掩码标签:按 15% 概率遮盖字符,其中 80% 替换为
[MASK]、10% 保留原字、10% 替换为随机词,同时记录被遮盖位置与原始标签; - 匹配 N-gram:遍历长度 2~7 的所有字符片段,在
ngram_dict词表中查询是否存在对应 N-gram,记录其 id、起始位置、长度与所属分段。
上述逻辑分别对应create_instances_from_document与create_masked_lm_predictions两个核心函数(examples/create_pre_train_data.py)。
六、输出文件说明:epoch_N.json 与 metrics
脚本会在--output_dir下生成两类文件:
epoch_0.json、epoch_1.json……:每行一条 JSON 训练样本,包含tokens、segment_ids、is_random_next、masked_lm_positions、masked_lm_labels以及ngram_ids、ngram_positions、ngram_lengths、ngram_tuples等 N-gram 字段;epoch_0_metrics.json……:记录该轮样本总数num_training_examples、max_seq_len等元信息,供训练脚本校验数据完整性。
run_pre_train.py会按epoch % num_data_epochs循环读取这些文件,所以生成 3 轮数据、训练 20 个 epoch 也是可行的(examples/run_pre_train.py)。
七、完整实战命令:一键生成预训练数据
以中文维基百科语料为例,一条可直接运行的命令如下:
python create_pre_train_data.py \ --train_corpus /data/zhwiki/zhwiki.txt \ --output_dir /data/zhwiki/pregenerated_data \ --bert_model /data/bert/bert-base-chinese \ --do_lower_case \ --do_whole_word_mask \ --reduce_memory \ --epochs_to_generate 3 \ --max_seq_len 128 \ --masked_lm_prob 0.15 \ --max_predictions_per_seq 20命令执行时会先显示Loading Dataset进度条加载语料,随后逐文档生成实例,最终在输出目录中得到 3 组 JSON 文件。若语料为空行分隔不正确,脚本会提示 "No document breaks were found",按上文第三部分修正即可。
八、常见问题与避坑指南
- 内存爆炸:语料达到 GB 级时,加上
--reduce_memory,文档会暂存到磁盘的 shelve 数据库中,而不是全部驻留内存; - ngram.txt 找不到:确认 N-gram 词表与
bert_model在同一目录,文件名必须是ngram.txt,格式为词,频次每行一条; - 样本数对不上:
metrics.json中的样本数是训练端断言依据,若手动删改 JSON 会导致run_pre_train.py断言失败; - 中文大小写:中文语料建议统一开启
--do_lower_case,与中文 BERT 分词器保持一致。
九、小结
create_pre_train_data.py是 ZEN 预训练流水线的第一环,它把普通中文语料"翻译"成模型能吃的字符级与 N-gram 级双重训练信号。掌握它的参数与输出格式,你就能顺利衔接run_pre_train.py,从零训练属于自己的中文 ZEN 模型。准备好语料,现在就动手试试吧!🚀
【免费下载链接】ZENA BERT-based Chinese Text Encoder Enhanced by N-gram Representations项目地址: https://gitcode.com/gh_mirrors/zen10/ZEN
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考