首页
/ TF-Models Model Garden NLP 预训练模型指南:BERT/ALBERT/ELECTRA 检查点与 TF-Hub 微调实战

TF-Models Model Garden NLP 预训练模型指南:BERT/ALBERT/ELECTRA 检查点与 TF-Hub 微调实战

2026-09-05 10:45:27作者:凌朦慧Richard

Model Garden(TensorFlow Model Garden 仓库)的 official/nlp 提供了一批可直接用于微调的 NLP 预训练模型资产,核心文档即 official/nlp/docs/pretrained_models.md。本文围绕该文档展开:先讲清“从检查点初始化(task.init_checkpoint)”与“加载 TF-Hub/Kaggle SavedModel(task.hub_module_url)”两种加载方式及其在训练驱动 official/nlp/train.py 中的源码机制,再完整给出 BERT、BERT 变体、ALBERT、ELECTRA 的模型清单与训练数据信息,最后结合仓库内真实实验配置(glue_mnli_matched.yamlsquad_v1.yamlbert_en_uncased_base.yaml)给出可复制运行的端到端微调命令。读完本文,你可以把任意一个官方预训练模型接入本仓库的训练框架完成 GLUE/SQuAD 等下游任务微调。

许可与数据集声明

在介绍模型之前,原文档给出两条重要的合规性声明,使用模型前务必注意:

  1. 数据集限制:这些检查点(Checkpoints)基于公开数据集训练,部分数据集存在非商用等使用限制。使用前请审阅第三方数据集提供方的条款与条件。
  2. 模型许可:检查点以 Apache 2.0 许可发布(对应仓库 LICENSE)。
  3. 数据集归属:文档中链接到的数据集不归 Google 所有或分发,均由第三方提供,同样需遵守第三方条款。

如何加载预训练模型

方式一:从检查点初始化(task.init_checkpoint

原文档首先说明:TF-Hub/Kaggle SavedModel 是官方推荐的模型分发方式,因为它自包含(self-contained),微调任务应优先考虑走 TF-Hub/Kaggle。如果仍需使用本地/GCS 上的检查点文件,可以直接使用 NLP 训练库(即 official/nlp/train.py),在启动任务时通过 params_override 指定检查点路径:

python3 train.py \
 --params_override=task.init_checkpoint=PATH_TO_INIT_CKPT

其中 PATH_TO_INIT_CKPT 通常指向 .tar.gz 解压目录中的 bert_model.ckpt(例如 SQuAD 微调示例中的 $BERT_DIR/bert_model.ckpt)。

方式二:加载 TF-Hub/Kaggle SavedModel(task.hub_module_url

像 SQuAD(阅读理解)和 GLUE(句子预测)这类微调任务支持直接从 TF-Hub/Kaggle 加载模型。这些内置任务专门支持 task.hub_module_url 参数:把 --params_override=task.init_checkpoint=... 替换为 --params_override=task.hub_module_url=TF_HUB_URL 即可,例如:

python3 train.py \
 --params_override=task.hub_module_url=https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/3

源码机制:两种方式的实现路径与互斥约束

结合仓库源码,可以确认两种方式在任务构建阶段的真实行为:

  • 互斥校验:在 official/nlp/tasks/sentence_prediction.pySentencePredictionConfig 中定义了 init_checkpointhub_module_url 两个字段,注释明确“At most one of init_checkpoint and hub_module_url can be specified”;build_model()(第 77–85 行)中若两者同时非空则抛出 ValueError。SQuAD 任务 official/nlp/tasks/question_answering.py 做了同样的双重校验(训练模型与评估模型各校验一次)。
  • Hub 加载路径:指定 hub_module_url 时,任务调用 official/nlp/tasks/utils.py 中的 get_encoder_from_hub()。该函数构造三个具名输入(input_word_idsinput_maskinput_type_ids),用 hub.KerasLayer(hub_model_path, trainable=True) 包装 TF-Hub 模块并输出其 output_dict,因此编码器权重是可训练的,后续会随微调一起更新。
  • 配置加载路径:未指定 hub_module_url 时,任务通过 encoders.build_encoder(self.task_config.model.encoder) 依据 YAML 中的编码器配置随机初始化网络。编码器配置类集中在 official/nlp/configs/encoders.pyEncoderConfig(one-of 结构,支持 bertalbertmobilebertxlnetfunnelbigbirdfnet 等),工厂函数 build_encoder(第 364 行起)按 config.type 分派到对应的 networks.*Encoder 实现。
  • 架构以 Hub 为准official/nlp/docs/train.md 特别说明——当从预训练模型初始化时,预训练模型自身的编码器架构会被采用,而你在配置文件里写的编码器架构(如 bert_en_uncased_base.yaml)会被忽略。因此使用 hub_module_url 时,模型结构无需与 YAML 严格一致。

另外,train.py 主流程(official/nlp/train.py)要求 --experiment--mode--model_dir 三个必填 FLAG,配置解析由 train_utils.parse_configuration(FLAGS) 完成,支持 YAML 文件与 params_override 的多级覆盖;--mode=continuous_train_and_eval 时还会调用 official/nlp/continuous_finetune_lib.py,其中要求 task.init_checkpoint 必须是一个目录(轮询等待其出现后加载最新检查点)。

BERT 预训练模型

文档说明:这里的 BERT 模型是 BERT 作者公开发布的预训练模型。仓库同时提供检查点tf.hub 模块两种微调用预训练资产,均为 TF 2.x 兼容版本,由 TF 1.x 官方 BERT 仓库(google-research/bert)发布的检查点转换而来,以保证与 BERT 论文一致。

检查点清单

模型 配置 训练数据 检查点与词表(GCS 对象,tf_model_garden/nlp/bert/v3/) TF-Hub/Kaggle SavedModel
BERT-base uncased 英文 uncased_L-12_H-768_A-12 Wiki + Books uncased_L-12_H-768_A-12.tar.gz tensorflow/bert_en_uncased_L-12_H-768_A-12(如 /3、/4 版本)
BERT-base cased 英文 cased_L-12_H-768_A-12 Wiki + Books cased_L-12_H-768_A-12.tar.gz tensorflow/bert_en_cased_L-12_H-768_A-12
BERT-large uncased 英文 uncased_L-24_H-1024_A-16 Wiki + Books uncased_L-24_H-1024_A-16.tar.gz tensorflow/bert_en_uncased_L-24_H-1024_A-16
BERT-large cased 英文 cased_L-24_H-1024_A-16 Wiki + Books cased_L-24_H-1024_A-16.tar.gz tensorflow/bert_en_cased_L-24_H-1024_A-16
BERT-large uncased(整词掩码) wwm_uncased_L-24_H-1024_A-16 Wiki + Books wwm_uncased_L-24_H-1024_A-16.tar.gz tensorflow/bert_en_wwm_uncased_L-24_H-1024_A-16
BERT-large cased(整词掩码) wwm_cased_L-24_H-1024_A-16 Wiki + Books wwm_cased_L-24_H-1024_A-16.tar.gz tensorflow/bert_en_wwm_cased_L-24_H-1024_A-16
BERT-base 多语言 multi_cased_L-12_H-768_A-12 Wiki + Books multi_cased_L-12_H-768_A-12.tar.gz tensorflow/bert_multi_cased_L-12_H-768_A-12
BERT-base 中文 chinese_L-12_H-768_A-12 Wiki + Books chinese_L-12_H-768_A-12.tar.gz tensorflow/bert_zh_L-12_H-768_A-12

更完整的 BERT 资产可在 TF-Hub 的 BERT 集合中查看(集合路径 google/collections/bert/1)。

配置命名遵循 L-层数_H-隐层维度_A-注意力头数 约定:uncased_L-12_H-768_A-12 即 12 层、768 维隐层、12 头。这与 official/nlp/configs/models/bert_en_uncased_base.yamlbert 编码器配置一一对应(num_layers: 12hidden_size: 768num_attention_heads: 12vocab_size: 30522max_position_embeddings: 512 等),也即上文 encoders.pyBertEncoderConfig 的默认值。

BERT 变体

除标准 BERT 外,文档还列出了在网络架构训练方法上做了变体的预训练模型(原文注明这些模型在下游任务上取得更高的准确率):

模型 配置 训练数据 TF-Hub SavedModel 说明
BERT-base talking heads + ggelu uncased_L-12_H-768_A-12 Wiki + Books tensorflow/talkheads_ggelu_bert_en_base BERT-base,采用 talking heads attention(arXiv:2003.02436)与 gated GeLU(arXiv:2002.05202)训练
BERT-large talking heads + ggelu uncased_L-24_H-1024_A-16 Wiki + Books tensorflow/talkheads_ggelu_bert_en_large BERT-large,同样采用 talking heads attention 与 gated GeLU
LAMBERT-large uncased 英文 uncased_L-24_H-1024_A-16 Wiki + Books tensorflow/lambert_en_uncased_L-24_H-1024_A-16 用 LAMB 优化器并借鉴 RoBERTa 技术训练的 BERT

ALBERT 预训练模型

ALBERT 的完整学术描述见其论文(arXiv:1909.11942)。与 BERT 相同,仓库同时提供检查点与 tf.hub 模块两种微调资产:TF 2.x 兼容,由 TF 1.x 官方 ALBERT 仓库(google-research/albert)发布的 ALBERT v2 检查点转换而来,与 ALBERT 论文保持一致。原文档特别强调:当前发布的检查点与 TF 1.x 官方 ALBERT 仓库完全相同

模型 训练数据 检查点与词表(GCS 对象,tf_model_garden/nlp/albert/) TF-Hub SavedModel
ALBERT-base 英文 Wiki + Books albert_base.tar.gz tensorflow/albert_en_base(/3 版本)
ALBERT-large 英文 Wiki + Books albert_large.tar.gz tensorflow/albert_en_large(/3 版本)
ALBERT-xlarge 英文 Wiki + Books albert_xlarge.tar.gz tensorflow/albert_en_xlarge(/3 版本)
ALBERT-xxlarge 英文 Wiki + Books albert_xxlarge.tar.gz tensorflow/albert_en_xxlarge(/3 版本)

对应的编码器 YAML 见 official/nlp/configs/models/albert_base.yamltype: albertvocab_size: 30000hidden_size: 768num_layers: 12 等),EncoderConfig.albert 分支会将其构建为 official/nlp/modeling/networks 中的 AlbertEncoder

ELECTRA 预训练模型

ELECTRA(Efficiently Learning an Encoder that Classifies Token Replacements Accurately,arXiv:2003.10555)是一种高效的预训练方法,原文档对其原理作了如下说明:

  • ELECTRA 包含两个 Transformer 模型:**generator(生成器)**与 discriminator(判别器)
  • 给定一条带掩码的序列,generator 将掩码位置的词替换为随机生成的词;discriminator 接收被破坏的句子,预测每个词是否被 generator 替换过。
  • 预训练阶段联合学习两个模型:generator 用掩码语言建模(MLM)任务训练,discriminator 用替换词检测(RTD,Replaced Token Detection)任务训练。
  • 微调阶段丢弃 generator,仅使用 discriminator 完成下游任务(如 GLUE、SQuAD)。
模型 训练数据 检查点与词表(GCS 对象,tf_model_garden/nlp/electra/)
ELECTRA-small 英文 Wiki + Books small.tar.gz(词表与 BERT uncased 英文相同)
ELECTRA-base 英文 Wiki + Books base.tar.gz(词表与 BERT uncased 英文相同)

原文档注明:这些检查点是用本仓库中的 Electra 代码重新训练的,仓库内也配套了 official/nlp/tasks/electra_task.py 与对应的任务配置(official/nlp/configs/electra.py)。

与训练配置的对应关系

上面的模型清单要与 train.py 的实验注册体系配合使用。实验名到配置的映射在 official/nlp/configs/finetuning_experiments.py 中注册,例如 bert/sentence_prediction(GLUE 类)、bert/squadbert/tagging 等;典型实验 YAML 位于 official/nlp/configs/experiments/

  • official/nlp/configs/experiments/glue_mnli_matched.yaml:句子预测(三分类)实验,其中 task.hub_module_url: ''task.init_checkpoint: '' 均为空字符串占位——也就是说这两个字段就是为本文所述的两种预训练加载方式预留的覆盖点;默认优化器为 AdamW(学习率 3e-5,多项式衰减,warmup 约 10% 步数),train_steps: 36813(392,702 条样本 × 3 epochs),metric_type: accuracy,按 cls_accuracy 导出最佳检查点。
  • official/nlp/configs/experiments/squad_v1.yaml:SQuAD 实验,同样预留空的 task.hub_module_url / task.init_checkpointseq_length: 384max_answer_length: 30n_best_size: 20,按 final_f1 选择最佳检查点,version_2_with_negative 控制 SQuAD 1.1/2.0。

--config_file 可传多个 YAML(后者优先),再叠加 --params_override 即可只改“模型来源”与“数据路径”,例如 official/nlp/MODEL_GARDEN.md 中给出的官方组合:MNLI-matched 微调用 task.hub_module_url=tensorflow/bert_en_uncased_L-12_H-768_A-12/4,SQuAD v1.1 微调用 task.hub_module_url=tensorflow/albert_en_base/3,也提供了用 GCS 检查点路径 gs://tf_model_garden/nlp/bert/uncased_L-12_H-768_A-12/bert_model.ckpt 作为 task.init_checkpoint 的对照写法。

端到端微调实战示例

以下两个完整流程来自仓库的 official/nlp/docs/train.md,与本文两种加载方式一一对应,可直接复制运行(数据需先下载并预处理为 tf_record,预处理脚本为 official/nlp/data/create_finetuning_data.py)。

示例一:用预训练检查点微调 SQuAD(init_checkpoint 路径)

在 Cloud TPU 上用预训练 BERT 检查点微调 SQuAD v1.1。先用检查点包中的 vocab.txt 生成 tf_record 训练数据:

export SQUAD_DIR=~/squad
export BERT_DIR=~/uncased_L-12_H-768_A-12   # BERT 检查点 tar.gz 解压目录
export OUTPUT_DATA_DIR=gs://some_bucket/datasets

python3 create_finetuning_data.py \
 --squad_data_file=${SQUAD_DIR}/train-v1.1.json \
 --vocab_file=${BERT_DIR}/vocab.txt \
 --train_data_output_path=${OUTPUT_DATA_DIR}/train.tf_record \
 --meta_data_file_path=${OUTPUT_DATA_DIR}/squad_meta_data \
 --fine_tuning_task_type=squad --max_seq_length=384

注:若要处理 SQuAD 2.0,需追加 FLAG --version_2_with_negative=True

然后启动训练与评估,注意 task.init_checkpoint 指向 bert_model.ckpt 文件:

export SQUAD_DIR=~/squad
export INPUT_DATA_DIR=gs://some_bucket/datasets
export OUTPUT_DIR=gs://some_bucket/my_output_dir
export BERT_DIR=~/uncased_L-12_H-768_A-12

export PARAMS=task.train_data.input_path=$INPUT_DATA_DIR/train.tf_record
export PARAMS=$PARAMS,task.validation_data.input_path=$SQUAD_DIR/dev-v1.1.json
export PARAMS=$PARAMS,task.validation_data.vocab_file=$BERT_DIR/vocab.txt
export PARAMS=$PARAMS,task.init_checkpoint=$BERT_DIR/bert_model.ckpt
export PARAMS=$PARAMS,runtime.distribution_strategy=tpu

python3 train.py \
 --experiment=bert/squad \
 --mode=train_and_eval \
 --model_dir=$OUTPUT_DIR \
 --config_file=configs/models/bert_en_uncased_base.yaml \
 --config_file=configs/experiments/squad_v1.1.yaml \
 --tpu=${TPU_NAME} \
 --params_override=$PARAMS

其中 train_data 是预处理后的 tf_record,而 validation_data 是原始 JSON(由数据加载器在线处理)。

示例二:从 TF-Hub 加载 BERT 微调 GLUE/MNLI(hub_module_url 路径)

同样先预处理 MNLI 数据(--fine_tuning_task_type=classification --max_seq_length=128 --classification_task_name=MNLI),然后:

export INPUT_DATA_DIR=gs://some_bucket/datasets
export OUTPUT_DIR=gs://some_bucket/my_output_dir
# TF-Hub 官方预训练 BERT 地址
export BERT_HUB_URL=https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/3

export PARAMS=task.train_data.input_path=$INPUT_DATA_DIR/mnli_train.tf_record
export PARAMS=$PARAMS,task.validation_data.input_path=$INPUT_DATA_DIR/mnli_eval.tf_record
export PARAMS=$PARAMS,task.hub_module_url=$BERT_HUB_URL
export PARAMS=$PARAMS,runtime.distribution_strategy=tpu

python3 train.py \
 --experiment=bert/sentence_prediction \
 --mode=train_and_eval \
 --model_dir=$OUTPUT_DIR \
 --config_file=configs/models/bert_en_uncased_base.yaml \
 --config_file=configs/experiments/glue_mnli_matched.yaml \
 --tfhub_cache_dir=$OUTPUT_DIR/hub_cache \
 --tpu=${TPU_NAME} \
 --params_override=$PARAMS

这里 --tfhub_cache_dir 用于缓存 TF-Hub 模块;训练进度可在控制台观察,输出模型位于 $OUTPUT_DIR。在 GPU 上运行时,把 runtime.distribution_strategy 设为 mirrored 并去掉 --tpu 即可。

小结

  • 本仓库 NLP 微调入口是 official/nlp/train.py,预训练模型只通过两个参数接入:task.init_checkpoint(检查点文件)与 task.hub_module_url(TF-Hub/Kaggle SavedModel),二者最多指定其一(源码中有显式 ValueError 校验)。
  • 官方推荐优先使用 TF-Hub/Kaggle SavedModel(自包含、架构自动跟随 Hub 模型);使用检查点时则需自备解压目录中的 vocab.txtbert_model.ckpt,并注意模型 YAML 中的编码器架构在 Hub 模式下会被忽略。
  • 可加载资产覆盖 BERT(含 WWM、多语言、中文及 talking heads/ggelu、LAMBERT 变体)、ALBERT(base 至 xxlarge,共 4 档)与 ELECTRA(small、base,均为 Wiki + Books 训练,词表同 BERT uncased)。
  • 实验 YAML(glue_mnli_matched.yamlsquad_v1.yaml 等)中的 hub_module_url/init_checkpoint 空值占位字段,就是 --params_override 的覆盖点,模型清单与实验配置由此闭环。

更多 FAQ 见 official/nlp/docs/faq.md,完整训练流程(含 Cloud TPU 环境搭建)见 official/nlp/docs/train.md

登录后查看全文
热门项目推荐
相关项目推荐