mode/models TF-NLP 预训练指南:从零预训练 BERT 与基于 TFDS 语料的掩码语言建模
2026/9/6 20:08:43 网站建设 项目流程

mode/models TF-NLP 预训练指南:从零预训练 BERT 与基于 TFDS 语料的掩码语言建模

【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models

本文基于 TensorFlow Models 仓库(mode/models)的 TF-NLP 官方文档 official/nlp/docs/pretrain.md 整理,系统讲解在 TPU/GPU 上运行 BERT 预训练实验的完整流程:先用原始 Wiki+Books 语料离线生成 TFRecord,再启动bert/pretraining实验;以及直接使用tensorflow_datasets的 Wikipedia 语料、以tf.text在线做 tokenize 与 masking 的bert/text_wiki_pretraining实验。读完后你可以复制文档中的命令与 YAML 配置独立完成预训练任务,并理解数据加载、损失计算与训练驱动(train.py)的底层实现。

一、两条预训练管线与实验注册机制

TF-NLP 的预训练入口是 official/nlp/train.py,它通过--experiment参数选择注册在exp_factory中的实验配置。预训练相关的实验定义集中在 official/nlp/configs/pretraining_experiments.py 中,包括:

  • bert/pretraining(L53-L67):读取离线预处理的 TFRecord(Wiki+Books 场景),使用BertPretrainDataConfig数据配置;
  • bert/pretraining_dynamic(L70-L87):动态长度输入序列版本,源码 docstring 说明其需要 TPU 搭配 tf.data service 的 round-robin 行为;
  • bert/text_wiki_pretraining(L90-L115):直接从 TFDS 的wikipedia/20201201.entrainsplit)读取文本,用tf.text在线预处理,仅使用英文 Wikipedia 语料;
  • electra/pretraining(L118-L135):ELECTRA 判别器式预训练,同样复用离线 TFRecord 数据。

这些实验的默认 trainer 配置(_TRAINER,L27-L50)为:adamw优化器(weight_decay_rate=0.01,对LayerNorm/layer_norm/bias排除权重衰减)、多项式学习率衰减(initial_learning_rate=1e-4衰减到0.0)、多项式 warmup,默认train_steps=1000000

训练驱动本身的调用链可以在 official/nlp/train.py 中看到:main()先解析 gin 配置,再调用train_utils.parse_configuration(FLAGS)完成「默认值 + 多个--config_fileYAML +--params_override字符串」的三层合并(实现在 official/core/train_utils.py 的parse_params_override/parse_configuration),随后进入_run_experiment_with_preemption_recovery——它会在 TPU 被抢占时从最近 checkpoint 自动重启训练。--experiment--mode--model_dir三个 flag 在 train.py L113 被标记为必填。

二、管线一:用原始 Wiki+Books 语料从零预训练 BERT

这条管线复现原始 BERT 论文的预训练配方:以 Wikipedia 和 Books 语料为输入,语料准备细节可参考原始 BERT 仓库中create_pretraining_data.py的说明(脚本在 official/nlp/data/create_pretraining_data.py,分支自 BERT 研究仓库并适配了 TF2 符号与 Python 3)。预训练配方是通用的,同样可以套用到你自己的语料上。

2.1 准备原始文本的格式要求

create_pretraining_data.py对输入文本有两点硬性约定,见 create_training_instances 的输入格式注释:

  1. 每行一个句子(理想情况是真正的句子而非整段文本,因为 next sentence prediction 任务依赖句子边界);
  2. 文档之间用空行分隔(文档边界用于保证 next sentence 任务不会跨文档采样)。

脚本按random_seed洗牌文档,并在dupe_factor次循环中为每份文档生成带不同 mask 的训练实例。

2.2 运行数据生成脚本:完整命令与参数说明

运行预训练脚本需要一个输入目录、输出目录以及词表文件。注意max_seq_length必须与你之后启动预训练时指定的序列长度一致。官方文档给出的标准命令:

export WORKING_DIR='local disk or cloud location' export BERT_DIR='local disk or cloud location' python models/official/nlp/data/create_pretraining_data.py \ --input_file=$WORKING_DIR/input/input.txt \ --output_file=$WORKING_DIR/output/tf_examples.tfrecord \ --vocab_file=$BERT_DIR/wwm_uncased_L-24_H-1024_A-16/vocab.txt \ --do_lower_case=True \ --max_seq_length=512 \ --max_predictions_per_seq=76 \ --masked_lm_prob=0.15 \ --random_seed=12345 \ --dupe_factor=5

(命令中的models/official/...为 tensorflow/models 仓库根下的相对路径;在本仓库中脚本即 official/nlp/data/create_pretraining_data.py。)

脚本的全部 flag 及默认值可从源码 L30-L99 确认:

Flag默认值说明
--input_file必填输入原始文本文件(或逗号分隔的多个文件/通配符,内部用tf.io.gfile.glob展开)
--output_file必填输出 TFRecord 文件(或逗号分隔的多个文件)
--tokenizationWordPiece分词器类型,可选WordPieceSentencePiece;标准 BERT 用 WordPiece,ALBERT 用 SentencePiece
--vocab_fileWordPiece 分词的词表文件
--sp_model_fileSentencePiece 分词的模型文件路径
--do_lower_caseTrue是否转小写:uncased 模型为 True,cased 模型为 False
--do_whole_word_maskFalse是否用整词 mask 代替逐 token mask
--max_ngram_sizeNone连续整词 n-gram mask 的最大长度(配合 zipf 加权偏好短 n-gram),需同时设置--do_whole_word_mask=True
--gzip_compressFalse是否输出 GZIP 压缩的 TFRecord
--use_v2_feature_namesFalse是否使用与模型输入一致的 v2 特征名(input_word_ids/input_type_ids
--max_seq_length128最大序列长度,须与预训练配置一致
--max_predictions_per_seq20每条序列的掩码 LM 预测数
--random_seed12345数据生成随机种子
--dupe_factor10输入数据复制次数(每次带不同 mask)
--masked_lm_prob0.15掩码语言建模概率
--short_seq_prob0.1生成短于最大长度序列的概率

2.3 掩码策略的源码细节

create_masked_lm_predictions(L587-L635)实现了经典的 80/10/10 策略:对被选中的预测位置,80% 概率整段替换为[MASK],10% 概率保留原词,10% 概率替换为词表中的随机词(L610-L623)。若启用--do_whole_word_mask,会先用_tokens_to_grams把 WordPiece 子词还原成整词区间,再由_masking_ngrams(L451-L554)按 zipf 加权(weight(n)=1/n,偏好短 n-gram)采样连续整词组。每个实例最终写入的特征包括input_ids(或 v2 命名)、input_masksegment_idsmasked_lm_positionsmasked_lm_idsmasked_lm_weightsnext_sentence_labels,见 write_instance_to_example_files。

A/B 句对构造上,脚本按句子切分文档,50% 概率让 B 句来自随机文档(is_random_next=True),否则取真实下一段;超过max_seq_length - 3(扣除[CLS]/[SEP]/[SEP])时随机从头部或尾部截断(truncate_seq_pair)。

2.4 更新实验 YAML:数据路径、分片与超参

生成 TFRecord 后,需要更新 YAML 实验配置,例如 official/nlp/configs/experiments/wiki_books_pretrain.yaml,填入你的数据路径,并把掩码相关超参对齐到数据生成时的设定。当数据有多个分片(shard)时,可以用*通配符一次包含多个文件。该文件的完整内容:

task: init_checkpoint: '' model: cls_heads: [{activation: tanh, cls_token_idx: 0, dropout_rate: 0.1, inner_dim: 768, name: next_sentence, num_classes: 2}] train_data: drop_remainder: true global_batch_size: 512 input_path: '[Your processed wiki data path]*,[Your processed books data path]*' is_training: true max_predictions_per_seq: 76 seq_length: 512 use_next_sentence_label: true use_position_id: false use_v2_feature_names: true validation_data: drop_remainder: false global_batch_size: 512 input_path: '[Your processed wiki data path]-00000-of-00500,[Your processed books data path]-00000-of-00500' is_training: false max_predictions_per_seq: 76 seq_length: 512 use_next_sentence_label: true use_position_id: false use_v2_feature_names: true trainer: checkpoint_interval: 20000 max_to_keep: 5 optimizer_config: learning_rate: polynomial: cycle: false decay_steps: 1000000 end_learning_rate: 0.0 initial_learning_rate: 0.0001 power: 1.0 type: polynomial optimizer: type: adamw warmup: polynomial: power: 1 warmup_steps: 10000 type: polynomial steps_per_loop: 1000 summary_interval: 1000 train_steps: 1000000 validation_interval: 1000 validation_steps: 64

这些train_data字段与BertPretrainDataConfig(official/nlp/data/pretrain_dataloader.py L30-L46)一一对应:input_path支持逗号分隔多路径、seq_length=512max_predictions_per_seq=76必须与数据生成时一致、use_v2_feature_names=true表示 TFRecord 使用input_word_ids/input_type_ids作为键(与 Keras 模型输入名保持一致)。trainer 侧:100 万步、学习率 1e-4 多项式衰减到 0、1 万步 warmup、训练批大小 512、每 1000 步验证 64 步、每 20000 步存 checkpoint 且最多保留 5 份。

调整不同 BERT 规模:若训练不同大小的 BERT,需要把模型配置中的分类头内维改到与隐藏层维度一致:

model: cls_heads: [{activation: tanh, cls_token_idx: 0, dropout_rate: 0.1, inner_dim: 768, name: next_sentence, num_classes: 2}]

即以inner_dim匹配 encoder 的hidden_size。以默认模型配置 official/nlp/configs/models/bert_en_uncased_base.yaml 为例,它是 12 层、12 头、hidden_size: 768intermediate_size: 3072max_position_embeddings: 512vocab_size: 30522的 BERT base,cls_headsinner_dim: 768正与之匹配。

2.5 启动训练与评估任务

随后启动训练和评估作业,运行bert/pretraining实验(注册见 pretraining_experiments.py L53-L67)。文档给出的完整命令:

export OUTPUT_DIR=gs://some_bucket/my_output_dir export PARAMS=$PARAMS,runtime.distribution_strategy=tpu python3 train.py \ --experiment=bert/pretraining \ --mode=train_and_eval \ --model_dir=$OUTPUT_DIR \ --config_file=configs/models/bert_en_uncased_base.yaml \ --config_file=configs/experiments/wiki_books_pretrain.yaml \ --tpu=${TPU_NAME} \ --params_override=$PARAMS

要点说明:

  • --mode=train_and_eval会同时跑训练循环与周期性评估;train.py 在 mode 含train时还会把合并后的完整配置序列化到 model_dir,便于复现;
  • 两个--config_file按序叠加:先加载模型结构(bert_en_uncased_base.yaml),再用实验 YAML 覆盖数据与 trainer 字段;
  • --params_override是最后一层字符串覆盖,这里把runtime.distribution_strategy设为tpu;GPU 环境可改为mirrored
  • --tpu=${TPU_NAME}指向 TPU 地址(如grpc://...:8470),无 TPU 时可不传;
  • 若配置了runtime.mixed_precision_dtype,驱动会调用performance.set_mixed_precision_policy设置混合精度(GPU 用 float16,TPU 用 bfloat16,见 train.py L99-L107);
  • 驱动默认开启异步 checkpoint(--enable_async_checkpointing,train.py L41-L44),并在 TPU 抢占时自动从最近 checkpoint 恢复(L47-L83)。

2.6 数据加载与任务的底层实现

bert/pretraining实验使用BertPretrainDataLoader(pretrain_dataloader.py L49-L137)。其_name_to_featuresseq_length/max_predictions_per_seq定义固定长度特征;当use_next_sentence_label=True时解析next_sentence_labels,当use_position_id=True时解析position_ids。一个值得注意的细节是 L99-L105:tf.Example只支持 int64 而 TPU 只支持 int32,因此解码时把所有 int64 特征统一 cast 为 int32。

任务侧MaskedLMTask(official/nlp/tasks/masked_lm.py):build_modelencoders.build_encoder构建 encoder,并把cls_heads逐个实例化为ClassificationHead,最终包成BertPretrainerV2(L59-L74);build_losses(L76-L107)以masked_lm_weights加权计算 MLM 的稀疏交叉熵(除以权重和),若 labels 中存在next_sentence_labels则叠加二分类的 NSP 交叉熵,总损失为两者之和。这也解释了为什么 YAML 中use_next_sentence_label开关同时决定数据字段与损失项。

三、管线二:基于 TFDS 语料的 BERT MLM 预训练

第二个示例用tensorflow_datasets直接预训练 BERT MLM 模型,并用tf.text在 TPU 上做预处理(在线 tokenize + 在线 masking,无需离线生成 TFRecord)。注意:只使用英文 Wikipedia 语料。该实验的源码 docstring 还明确了一个限制(L90-L97):由于 next sentence 采样难以用 tf ops 精确匹配原始实现,这条管线无法完全复现 BERT 的完整训练设置——事实上实验的 TFDS 文本数据默认走纯 MLM。

实验使用BertPretrainTextDataConfig(official/nlp/data/pretrain_text_dataloader.py L30-L48),关键字段:tfds_name/tfds_split(实验默认wikipedia/20201201.entrainsplit)、vocab_file_path(YAML 中默认是占位符'Please provide the vocab file path.',需覆盖为真实词表路径)、masking_rate: 0.15use_whole_word_maskingdoc_batch_size: 8(NSP 场景下文档级批处理的文档数)。对应的实验 YAML 是 official/nlp/configs/experiments/wiki_tfds_pretrain.yaml:

task: init_checkpoint: '' model: cls_heads: [{activation: tanh, cls_token_idx: 0, dropout_rate: 0.1, inner_dim: 768, name: next_sentence, num_classes: 2}] train_data: drop_remainder: true global_batch_size: 512 is_training: true max_predictions_per_seq: 76 seq_length: 512 use_next_sentence_label: false use_whole_word_masking: true tfds_name: wikipedia/20201201.en tfds_split: train vocab_file_path: 'Please provide the vocab file path.' validation_data: drop_remainder: true global_batch_size: 32 is_training: false max_predictions_per_seq: 76 seq_length: 512 use_next_sentence_label: false use_whole_word_masking: true tfds_name: wikipedia/20201201.en tfds_split: train vocab_file_path: 'Please provide the vocab file path.' trainer: # 与 wiki_books_pretrain.yaml 相同的 adamw + polynomial 衰减 + warmup 配置 train_steps: 1000000 # ...(完整内容见文件)

与管线一的关键差异:use_next_sentence_label: false(纯 MLM,不产生next_sentence_labels)、新增use_whole_word_masking: true、验证集global_batch_size降为 32。

BertPretrainTextDataLoader的在线处理流程(L105-L202)全部由 tf.text/tf.data op 组成,可在 TPU 上高效执行:用词表构建StaticVocabularyTable并查出[CLS]/[SEP]/[MASK]的 id →tf_text.BertTokenizer分词(whole-word masking 时保留 wordpieces 维度)→WaterfallTrimmer截断到seq_length - 3tf_text.combine_segments拼接并生成 segment ids →tf_text.RandomItemSelectorselection_ratemasking_rate,排除 CLS/SEP)加MaskValuesChoosertf_text.mask_language_model动态生成掩码输入 →pad_model_inputs填充到固定形状并导出input_maskmasked_lm_weightsload()中还有一处 NSP 的约束:开启use_next_sentence_label时只支持单个 text 字段(L69-L71)。

启动训练与评估作业(运行bert/text_wiki_pretraining实验,注册见 pretraining_experiments.py L90-L115)。文档给出的命令——用 FLAGS 覆盖配置,或直接编辑configs/experiments/wiki_tfds_pretrain.yaml对应字段均可:

export OUTPUT_DIR=gs://some_bucket/my_output_dir # 更多预训练 checkpoint 说明见仓库文档 official/nlp/docs/pretrained_models.md export BERT_DIR=~/cased_L-12_H-768_A-12 # Override the configurations by FLAGS. Alternatively, you can directly edit # `configs/experiments/wiki_tfds_pretrain.yaml` to specify corresponding fields. export PARAMS=$PARAMS,task.validation_data.vocab_file_path=$BERT_DIR/vocab.txt export PARAMS=$PARAMS,task.train_data.vocab_file_path=$BERT_DIR/vocab.txt export PARAMS=$PARAMS,runtime.distribution_strategy=tpu python3 train.py \ --experiment=bert/text_wiki_pretraining \ --mode=train_and_eval \ --model_dir=$OUTPUT_DIR \ --config_file=configs/experiments/wiki_tfds_pretrain.yaml \ --tpu=${TPU_NAME} \ --params_override=$PARAMS

这里--params_override用「层级路径=值」的形式把task.train_data.vocab_file_pathtask.validation_data.vocab_file_path指向cased_L-12_H-768_A-12/vocab.txt(注意:TFDS 管线用的是 cased 词表,若换 uncased 词表需同步调整数据预处理),并把分布式策略覆盖为 TPU。

四、两条管线的选型与配置核对清单

维度bert/pretrainingbert/text_wiki_pretraining
数据来源离线 TFRecord(Wiki+Books,需先跑create_pretraining_data.pyTFDSwikipedia/20201201.en(在线读取文本)
预处理离线:Python 端分词 + 80/10/10 masking,支持整词/n-gram mask在线:tf.text分词、截断、mask_language_model动态 masking
NSP支持(use_next_sentence_label: true默认关闭(源码说明采样难以完全匹配原实现)
分片input_path支持*通配多分片由 TFDS split 决定
数据配置类BertPretrainDataConfigBertPretrainTextDataConfig
实验 YAMLwiki_books_pretrain.yamlwiki_tfds_pretrain.yaml

无论走哪条管线,落盘前建议核对:seq_lengthmax_predictions_per_seq在三处(数据生成命令、实验 YAML 的train_data/validation_data、验证集配置)保持一致;input_path/vocab_file_path已替换为真实路径;模型 YAML 的hidden_sizecls_headsinner_dim匹配;runtime.distribution_strategy--tpu参数匹配目标硬件。仓库中 official/nlp/docs/ 目录下的pretrained_models.mdfaq.md等文档可作为预训练后续导出与部署的延伸参考。

【免费下载链接】modelsModels and examples built with TensorFlow项目地址: https://gitcode.com/GitHub_Trending/mode/models

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询