TF-Models Model Garden NLP 预训练模型指南:BERT/ALBERT/ELECTRA 检查点与 TF-Hub 微调实战
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.yaml、squad_v1.yaml、bert_en_uncased_base.yaml)给出可复制运行的端到端微调命令。读完本文,你可以把任意一个官方预训练模型接入本仓库的训练框架完成 GLUE/SQuAD 等下游任务微调。
许可与数据集声明
在介绍模型之前,原文档给出两条重要的合规性声明,使用模型前务必注意:
- 数据集限制:这些检查点(Checkpoints)基于公开数据集训练,部分数据集存在非商用等使用限制。使用前请审阅第三方数据集提供方的条款与条件。
- 模型许可:检查点以 Apache 2.0 许可发布(对应仓库 LICENSE)。
- 数据集归属:文档中链接到的数据集不归 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.py 的
SentencePredictionConfig中定义了init_checkpoint、hub_module_url两个字段,注释明确“At most one ofinit_checkpointandhub_module_urlcan 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_ids、input_mask、input_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.py 的EncoderConfig(one-of 结构,支持bert、albert、mobilebert、xlnet、funnel、bigbird、fnet等),工厂函数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.yaml 中 bert 编码器配置一一对应(num_layers: 12、hidden_size: 768、num_attention_heads: 12、vocab_size: 30522、max_position_embeddings: 512 等),也即上文 encoders.py 中 BertEncoderConfig 的默认值。
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.yaml(type: albert,vocab_size: 30000、hidden_size: 768、num_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/squad、bert/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_checkpoint;seq_length: 384、max_answer_length: 30、n_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.txt与bert_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.yaml、squad_v1.yaml等)中的hub_module_url/init_checkpoint空值占位字段,就是--params_override的覆盖点,模型清单与实验配置由此闭环。
更多 FAQ 见 official/nlp/docs/faq.md,完整训练流程(含 Cloud TPU 环境搭建)见 official/nlp/docs/train.md。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0623
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00