tensorflow/models TF-NLP:用 export_tfhub 将预训练 Transformer Encoder 导出为 TF Hub SavedModel 全指南
本文基于 TF-NLP 仓库的官方文档 tfhub.md 及其配套源码,系统讲解如何使用 official/nlp/tools/export_tfhub.py 工具,把 BERT、ALBERT 等预训练 Transformer 编码器导出为符合 TF Hub 通用 API 的 SavedModel。读完后,你将掌握:编码器模型(含 MLM 头)与预处理模型两类导出的完整命令行操作、所有关键参数(--export_type、--encoder_config_file、--model_checkpoint_path、--vocab_file/--sp_model_file、--do_lower_case、--copy_pooler_dense_to_encoder 等)的取值含义与底层实现,以及导出后 SavedModel 的调用接口与发布前验证方法。
TF Hub 上文本编码器的"双模型"体系
在 TF Hub 上,文本 Transformer 编码器以一对 SavedModel 的形式发布:
- 预处理模型(preprocessing model):应用一个基于固定词表的 tokenizer,并附加少量逻辑,把原始文本转换为 Transformer 输入(
input_word_ids/input_mask/input_type_ids); - 编码器模型(encoder model,简称 model):应用预训练好的 Transformer 编码器本体。
TF Hub 为这两类 SavedModel 定义了统一的 Common API(Common SavedModel APIs for Text),把具体的预处理逻辑和编码器架构选择都封装在这两套标准接口之下。本文介绍的 export_tfhub 工具导出的正是符合这套 API 的 SavedModel。
工具入口 export_tfhub.py 的文件头 docstring 明确说明:该工具创建可直接上传至 tfhub.dev 的 preprocessor 与 encoder SavedModel,并实现了 TF Hub 文本 Common API 中定义的 preprocessor 与 encoder 接口。
总体参数:--export_type 的三种取值
--export_type 是总开关,取值为 model、model_with_mlm、preprocessing 三选一(定义于 export_tfhub.py):
| 取值 | 导出内容 | 必需配置 |
|---|---|---|
model |
仅编码器本体 | 编码器配置 + 检查点 + 词表 |
model_with_mlm |
编码器 + 预训练 MLM 头 | 编码器配置 + 更严格的检查点 + 词表 |
preprocessing |
预处理(tokenizer + 打包)模型 | 词表或 SP 模型 |
其余参数的两组"互斥对"在入口处就会被强制校验:
--vocab_file与--sp_model_file必须恰好设置一个(WordPiece 词表 vs SentencePiece 模型文件);--encoder_config_file与--bert_config_file必须恰好设置一个(导出model/model_with_mlm时)。
这一校验逻辑见 export_tfhub.py 的 main()。
仅导出编码器(--export_type=model)
命令示例
python official/nlp/tools/export_tfhub.py \
--encoder_config_file=${BERT_DIR:?}/bert_encoder.yaml \
--model_checkpoint_path=${BERT_DIR:?}/bert_model.ckpt \
--vocab_file=${BERT_DIR:?}/vocab.txt \
--export_type=model \
--export_path=/tmp/bert_model
--encoder_config_file 与 --bert_config_file
--encoder_config_file指向一个 YAML 文件,它表示 encoders.py 中定义的encoders.EncoderConfig数据类,支持多种编码器类型(BERT、ALBERT 等)。仓库中就有一份可直接参考的样例配置 bert_en_uncased_base.yaml,其结构为:
task:
model:
encoder:
type: bert
bert:
attention_dropout_rate: 0.1
dropout_rate: 0.1
hidden_activation: gelu
hidden_size: 768
initializer_range: 0.02
intermediate_size: 3072
max_position_embeddings: 512
num_attention_heads: 12
num_layers: 12
type_vocab_size: 2
vocab_size: 30522
- 若只想导出 BERT,也可以改用
--bert_config_file指向旧版的bert_config.json(legacy 格式)。源码中该路径会走configs.BertConfig.from_json_file()解析,并据此构建networks.BertEncoder(见 export_tfhub_lib.py 的get_bert_encoder());而--encoder_config_file路径则通过encoders.build_encoder()构建编码器(见 export_tfhub_lib.py)。 - 如果模型定义涉及 GIN 配置,还必须设置
--gin_file与--gin_params,并与预训练时保持一致。入口处会先执行gin.parse_config_files_and_bindings(FLAGS.gin_file, FLAGS.gin_params)(export_tfhub.py)。
--model_checkpoint_path:检查点格式要求
该参数指向一个对象化(TF2)检查点。对 --export_type=model:
- 要求检查点能恢复到
tf.train.Checkpoint(encoder=encoder),其中encoder是由配置参数定义的那个编码器对象; - 由
BertPretrainerV2(bert_pretrainer.py)写出的检查点天然满足这一点; - 兼容期约定:旧版以
model=而非encoder=保存的检查点同样支持。
从源码可以确认,实际恢复时是同时以两个名字挂到同一个编码器对象上再恢复,从而兼容新旧两种命名(export_tfhub_lib.py):
checkpoint = tf.train.Checkpoint(
model=encoder, # Legacy checkpoints.
encoder=encoder)
checkpoint.restore(model_checkpoint_path).assert_existing_objects_matched()
导出模型的输入输出接口
导出的 SavedModel 接受 dict 输入、返回 dict 输出,实现了对 TF Hub "Text embeddings with Transformer encoders" Common API 的特定化:
encoder = hub.load(...)
encoder_inputs = dict(
input_word_ids=..., # Shape [batch, seq_length], dtype=int32
input_mask=..., # Shape [batch, seq_length], dtype=int32
input_type_ids=..., # Shape [batch, seq_length], dtype=int32
)
encoder_outputs = encoder(encoder_inputs)
assert encoder_outputs.keys() == {
"pooled_output", # Shape [batch_size, width], dtype=float32
"default", # Alias for "pooled_output" (aligns with other models)
"sequence_output", # Shape [batch_size, seq_length, width], dtype=float32
"encoder_outputs", # List of Tensors with outputs of all transformer layers
}
其中 "default" 别名是在导出代码中显式注入的——为了让 BERT 类模型与其他文本表征模型在 TF Hub 上可互换,pooled_output 被复制为 default(export_tfhub_lib.py)。
pooler 层的三种处理方式
编码器导出的 pooler 层权重从 --model_checkpoint_path 恢复。但要注意:与经典 BERT 不同,BertPretrainerV2 不训练编码器自带的 pooler 层。文档给出三种应对方式:
- 设置
--copy_pooler_dense_to_encoder,把传入BertPretrainerV2用于下一句预测(NSP)任务的ClassificationHead中的 pooling 层复制到编码器上。这模仿了经典 BERT 的做法,但不推荐用于新模型; - 不设置该 flag,直接导出编码器未训练、随机初始化的 pooler 层。按 2020 年时的社区经验(folklore),未训练的 pooler 往往比预训练过的 pooler 更容易被微调,因此这是默认行为;
- 不设置该 flag,在导出前自行初始化 pooling 层。例如 Google 于 2020 年 10 月发布的 "BERT Experts" 将其初始化为恒等映射,报告微调效果相当、且行为更可预测。
无论采用哪种方式,导出工具目前都要求编码器模型必须带有 pooled_output(无论是否训练过)。
--copy_pooler_dense_to_encoder 的实现见 export_tfhub_lib.py:它额外构造一个 tf.train.Checkpoint(**{"next_sentence.pooler_dense": encoder.pooler_layer}) 并从同一检查点恢复,即从 NSP 分类头的命名空间取出 dense 层权重。测试用例 export_tfhub_lib_test.py 验证了复制后的 pooled_output 与 NSP 头 dense 层输出一致、而与编码器原生 pooler 输出不同。
词表信息的附带导出
编码器模型本身不包含任何预处理逻辑。但为了方便自行做预处理的旧版用户,相关 tokenizer 信息会从 --vocab_file 或 --sp_model_file(二选一)以及 --do_lower_case 三个参数附着到导出产物上——这些值必须与导出预处理模型时完全一致。
导出后,根对象上会存储这些属性,可以这样读取(export_tfhub_lib.py 中将其写为 tf.saved_model.Asset 与不可训练的 tf.Variable):
encoder = hub.load(...)
# Gets the filename of the respective tf.saved_model.Asset object.
if hasattr(encoder, "vocab_file"):
print("Wordpiece vocab at", encoder.vocab_file.asset_path.numpy())
elif hasattr(encoder, "sp_model_file"):
print("SentencePiece model at", encoder.sp_model_file.asset_path.numpy())
# Gets the value of a scalar bool tf.Variable.
print("...using do_lower_case =", encoder.do_lower_case.numpy())
文档建议:新用户应忽略这些属性、直接使用预处理模型;这些属性是为遗留用户以及需要访问完整词表的高级用户保留的。
带 MLM 头导出(--export_type=model_with_mlm)
命令与检查点要求的差异
在上一节全部参数说明的基础上,把 --export_type 从 model 改为 model_with_mlm 即可:
python official/nlp/tools/export_tfhub.py \
--encoder_config_file=${BERT_DIR:?}/bert_encoder.yaml \
--model_checkpoint_path=${BERT_DIR:?}/bert_model.ckpt \
--vocab_file=${BERT_DIR:?}/vocab.txt \
--export_type=model_with_mlm \
--export_path=/tmp/bert_model
此时 --model_checkpoint_path 的要求更严格:检查点必须能恢复到 tf.train.Checkpoint(**BertPretrainerV2(...).checkpoint_items)。从源码看,checkpoint_items 包含 encoder、masked_lm,以及各分类头命名空间下的条目(bert_pretrainer.py),即除了编码器权重还要求 MLM 头的权重。因此并非所有 Transformer 编码器、所有预训练技术都能满足该要求——例如 ELECTRA 虽然使用 BERT 架构,但预训练时不使用 MLM 任务,就只能走 --export_type=model 路径(测试文件 export_tfhub_lib_test.py 的注释也印证了这一点:不带 .mlm 的导出"对 Electra 这类无 MLM 任务训练的模型仍然有用")。
.mlm 子对象接口
导出后根对象的调用方式与上节相同,额外多出一个可被调用的 mlm 子对象:
mlm_inputs = dict(
input_word_ids=..., # Shape [batch, seq_length], dtype=int32
input_mask=..., # Shape [batch, seq_length], dtype=int32
input_type_ids=..., # Shape [batch, seq_length], dtype=int32
masked_lm_positions=..., # Shape [batch, num_predictions], dtype=int32
)
mlm_outputs = encoder.mlm(mlm_inputs)
assert mlm_outputs.keys() == {
"pooled_output", # Shape [batch, width], dtype=float32
"sequence_output", # Shape [batch, seq_length, width], dtype=float32
"encoder_outputs", # List of Tensors with outputs of all transformer layers
"mlm_logits" # Shape [batch, num_predictions, vocab_size], dtype=float32
}
源码实现上,MLM 子模型是 BertPretrainerV2 包装后的 Keras 模型,导出时通过 core_model._auto_track_sub_layers = False 把 mlm 附加到根模型上,使 MLM 的额外权重不混入编码器本体的权重集合(export_tfhub_lib.py)。文档也提醒:这个额外子对象会带来适度的体积开销。测试 export_tfhub_lib_test.py 逐权重校验了 hub_layer.resolved_object.mlm 与 BertPretrainerV2 的权重一致,并验证 mlm_logits 形状为 [batch, num_predictions, vocab_size]。
从 TF1 BERT 检查点导出
用 TF1 原版 BERT 实现训练出的模型,可以先用 tf2_bert_encoder_checkpoint_converter.py 工具把检查点转换为对象化的 V2 检查点,该工具支持输出 encoder 或 pretrainer 两种命名空间(其 --converted_model 参数取值即为这两者),然后按上面的流程对转换后的检查点运行 export_tfhub。两点注意:
- 不要设置
--copy_pooler_dense_to_encoder,因为 pooler 层已经是转换后编码器的一部分; --vocab_file与--do_lower_case可以直接照搬 TF1 BERT 的取值。
导出预处理模型(--export_type=preprocessing)
如果 TF Hub 上已存在与你编码器需求完全匹配的预处理模型(相同 tokenizer、相同词表、相同 do_lower_case 归一化设置),可以跳过此步,直接复用已有模型的预处理模型。
命令示例
python official/nlp/tools/export_tfhub.py \
--vocab_file=${BERT_DIR:?}/vocab.txt \
--do_lower_case=True \
--export_type=preprocessing \
--export_path=/tmp/bert_preprocessing
参数说明
--vocab_file:与BertTokenizer配套使用的词表文件;若模型使用SentencepieceTokenizer,则改设--sp_model_file(两者互斥,见 export_tfhub.py 的 flag 定义)。--do_lower_case:控制文本归一化(与对应 tokenizer 类的行为一致,比单纯"压平大小写"略多)。若不设置,则遵循推断规则:--vocab_file路径中出现uncased时自动启用;设置--sp_model_file时无条件启用(模仿 BERT 与 ALBERT 的惯例)。源码中这一推断逻辑即 export_tfhub_lib.py 的get_do_lower_case()。程序化调用或拿不准时,建议显式设置--do_lower_case。--default_seq_length:当调用时省略seq_length参数时生效,默认值 128——因为 128 的倍数在 Cloud TPU 上表现最佳,而注意力计算成本随seq_length二次增长(flag 定义见 export_tfhub.py)。- GIN:如果预处理定义涉及 GIN 配置,同样需要设置
--gin_file/--gin_params且与预训练一致(撰写该文档时,代码中尚不存在这样的 GIN 可配置项)。 - TF 2.4 兼容性注意:面向 TensorFlow 2.4.x 公开版用户导出时,应设置
--experimental_disable_assert_in_preprocessing,以避免预处理在 TPU worker 的Dataset.map()中使用时发生致命的算子放置问题;该问题在 TF 2.3 与 TF 2.5+ 中不存在。其实现原理见 export_tfhub_lib.py:导出期间临时 monkey-patch 掉tf.Assert(替换为 no-op),导出完成后由_check_no_assert()解析saved_model.pb的 proto,扫描全局图与所有函数库中的 Assert 节点并自我校验(export_tfhub_lib.py)。
导出模型的调用接口
单段文本输入的最简调用方式:
preprocessor = hub.load(...)
text_input = ... # Shape [batch_size], dtype=tf.string
encoder_inputs = preprocessor(text_input, seq_length=seq_length)
assert encoder_inputs.keys() == {
"input_word_ids", # Shape [batch_size, seq_length], dtype=int32
"input_mask", # Shape [batch_size, seq_length], dtype=int32
"input_type_ids" # Shape [batch_size, seq_length], dtype=int32
}
除根调用外,导出的 SavedModel 完整实现了 TF Hub "Text embeddings with preprocessed inputs and Transformer encoders" 预处理 API 的子对象集合。从 export_tfhub_lib.py 的 create_preprocessing() 可见,这些子对象包括:
preprocess(根调用):一步完成 tokenize + 打包,适合单句输入;preprocess.tokenize:仅分词,输出 RaggedTensor 的 token 序列;preprocess.tokenize.get_special_tokens_dict:以无参tf.function形式暴露词表特殊 token 信息(vocab_size、padding_id、start_of_sequence_id、end_of_segment_id等);preprocess.tokenize_with_offsets:分词并同时给出起止偏移(由--tokenize_with_offsets控制导出与否);preprocess.bert_pack_inputs:把 1~2 段 RaggedToken 打包为定长输入,且支持调用时用seq_length=覆盖导出时的默认长度。
其中 bert_pack_inputs 之所以要用专门的包装器 BertPackInputsSavedModelWrapper,是因为直接保存 Keras 层会固定 RaggedTensor 的个数与 ragged rank,且超参数在保存后不可改;包装器在导出时对 4 种"rank × 段数"组合逐一预热了 concrete function(export_tfhub_lib.py),并允许 seq_length 作为调用期参数。
另外,源码中有一个值得注意的工程细节:导出时词表/SP 模型文件会被复制到一个临时目录再引用(_move_to_tmpdir,export_tfhub_lib.py),以避免本地绝对路径(含敏感路径分量)泄漏进公开的 SavedModel;测试 export_tfhub_lib_test.py 专门验证了 saved_model.pb 中不含原始目录名。
获取特殊 token 以支持 MLM
使用 encoder.mlm() 接口时,需要用用户代码对分词后的输入做随机掩码,所需的词表信息可以从预处理模型中统一获取(对 WordPiece 与 SentencePiece 两种 tokenizer 均适用):
special_tokens_dict = preprocess.tokenize.get_special_tokens_dict()
vocab_size = int(special_tokens_dict["vocab_size"])
padding_id = int(special_tokens_dict["padding_id"]) # [PAD] or <pad>
start_of_sequence_id = int(special_tokens_dict["start_of_sequence_id"]) # [CLS]
end_of_segment_id = int(special_tokens_dict["end_of_segment_id"]) # [SEP]
mask_id = int(special_tokens_dict["mask_id"]) # [MASK]
一个完整的"预处理 + MLM"端到端用法在测试 export_tfhub_lib_test.py 中:先 preprocess.tokenize 两段文本,用 tf.text.WaterfallTrimmer 裁剪、tf.text.combine_segments 合并、tf.text.mask_language_model 按策略随机掩码,再经 pad_model_inputs 定长化后调用 encoder.mlm(mlm_inputs),最终断言 mlm_logits 形状为 [batch_size, num_predictions, vocab_size]。同一测试还验证了 get_special_tokens_dict 中各特殊 token 的实际取值(WordPiece 布局为 [PAD]=0、[UNK]=1、[CLS]=2、[SEP]=3、[MASK]=4)。
源码实现速览:一次编码器导出发生了什么
把 export_tfhub.py 的 main() 与 export_tfhub_lib.py 串起来,一次 --export_type=model 的执行链为:
- 解析 GIN 配置,校验
vocab_file/sp_model_file互斥,推断do_lower_case; - 按配置构建编码器:
--bert_config_file走get_bert_encoder(),--encoder_config_file走encoders.build_encoder()(export_tfhub_lib.py); _create_model()把编码器的命名输入转为 dict 输入(只有 dict 形式在 SavedModel 恢复后才可接受 dict 调用),前向一次得到输出 dict,并注入default别名,包成tf_keras.Model(inputs=..., outputs=...)(export_tfhub_lib.py);- 构造双命名空间检查点(
model=/encoder=)恢复权重并断言匹配; - 视需要复制 NSP pooler dense 层;
- 把词表资产与
do_lower_case挂到根模型,core_model.save(export_path, include_optimizer=False, save_format="tf")完成落盘(export_tfhub_lib.py)。
--export_type=preprocessing 分支则走 export_preprocessing():在临时目录中复制词表资产、按需禁用 Assert 算子、create_preprocessing() 组装带子对象的模型并保存(export_tfhub_lib.py)。
用测试套件验证导出正确性
export_tfhub_lib_test.py 是验证导出行为的关键参照,覆盖三类用例:
ExportModelTest(不带 MLM):参数化覆盖 legacy BertConfig、ALBERT、BertEncoder、BertEncoderV2 四种编码器,用hub.KerasLayer(export_path, trainable=True)重新加载导出结果,逐权重比对、比对pooled_output/sequence_output/encoder_outputs输出、验证default == pooled_output、验证training=True时 dropout 生效(输出标准差 > 1e-3)、以及seq_length在形状推断中的传播;ExportModelWithMLMTest(带 MLM):在 BERT 与 ALBERT 两种配置下验证.mlm子对象存在、子对象权重与BertPretrainerV2逐一对应、mlm_logits输出与源模型一致,以及test_copy_pooler_dense_to_encoder专门验证 NSP pooler 复制行为;ExportPreprocessingTest:验证.tokenize、.tokenize_with_offsets、根调用、.bert_pack_inputs的 token 级精确输出(含input_type_ids的分段标记)、do_lower_case=False与自定义default_seq_length=10的组合、形状推断(XLA 友好的静态形状)、re-export(再 save/load 一轮后仍可独立加载工作,对应 TensorFlow issue 46456 的失败场景)、特殊 token 在 Estimator graph 模式下的获取,以及路径泄漏自检。
发布前的测试:微调对拍
文档最后强调:发布前务必用适当任务对导出的 SavedModel 做微调测试,并与等价的 Python 原生代码基线实验比较性能和精度。仓库中的 train.md(TF-NLP trainer 文档)提供了用 GLUE 基准(如 BERT on MNLI)跑微调的具体步骤,可作为这个"基线对拍"的参照实验。
实操要点清单
- 三类
--export_type各走不同代码分支,先确认自己导出什么再选参; --vocab_file与--sp_model_file、--bert_config_file与--encoder_config_file各自二选一,恰好设一个;- 编码器导出的检查点要求因是否带 MLM 而不同:
model只需Checkpoint(encoder=encoder)可恢复;model_with_mlm需要完整BertPretrainerV2.checkpoint_items(无 MLM 任务的模型如 ELECTRA 不能走这条路); --do_lower_case不显式设置时有基于词表文件名的推断规则,程序化场景建议显式指定;- 面向 TF 2.4 用户发布预处理模型时加
--experimental_disable_assert_in_preprocessing; - TF1 BERT 检查点先转换、转换后禁用
--copy_pooler_dense_to_encoder; - 发布前完成微调对拍验证。
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