首页
/ tensorflow/models TF-NLP:用 export_tfhub 将预训练 Transformer Encoder 导出为 TF Hub SavedModel 全指南

tensorflow/models TF-NLP:用 export_tfhub 将预训练 Transformer Encoder 导出为 TF Hub SavedModel 全指南

2026-09-05 21:51:56作者:鲍丁臣Ursa

本文基于 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 是总开关,取值为 modelmodel_with_mlmpreprocessing 三选一(定义于 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.pymain()

仅导出编码器(--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.pyget_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 是由配置参数定义的那个编码器对象;
  • BertPretrainerV2bert_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 被复制为 defaultexport_tfhub_lib.py)。

pooler 层的三种处理方式

编码器导出的 pooler 层权重从 --model_checkpoint_path 恢复。但要注意:与经典 BERT 不同,BertPretrainerV2 不训练编码器自带的 pooler 层。文档给出三种应对方式:

  1. 设置 --copy_pooler_dense_to_encoder,把传入 BertPretrainerV2 用于下一句预测(NSP)任务的 ClassificationHead 中的 pooling 层复制到编码器上。这模仿了经典 BERT 的做法,但不推荐用于新模型;
  2. 不设置该 flag,直接导出编码器未训练、随机初始化的 pooler 层。按 2020 年时的社区经验(folklore),未训练的 pooler 往往比预训练过的 pooler 更容易被微调,因此这是默认行为;
  3. 不设置该 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_typemodel 改为 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 包含 encodermasked_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 = Falsemlm 附加到根模型上,使 MLM 的额外权重不混入编码器本体的权重集合(export_tfhub_lib.py)。文档也提醒:这个额外子对象会带来适度的体积开销。测试 export_tfhub_lib_test.py 逐权重校验了 hub_layer.resolved_object.mlmBertPretrainerV2 的权重一致,并验证 mlm_logits 形状为 [batch, num_predictions, vocab_size]

从 TF1 BERT 检查点导出

用 TF1 原版 BERT 实现训练出的模型,可以先用 tf2_bert_encoder_checkpoint_converter.py 工具把检查点转换为对象化的 V2 检查点,该工具支持输出 encoderpretrainer 两种命名空间(其 --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.pyget_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.pycreate_preprocessing() 可见,这些子对象包括:

  • preprocess(根调用):一步完成 tokenize + 打包,适合单句输入;
  • preprocess.tokenize:仅分词,输出 RaggedTensor 的 token 序列;
  • preprocess.tokenize.get_special_tokens_dict:以无参 tf.function 形式暴露词表特殊 token 信息(vocab_sizepadding_idstart_of_sequence_idend_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_tmpdirexport_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.pymain()export_tfhub_lib.py 串起来,一次 --export_type=model 的执行链为:

  1. 解析 GIN 配置,校验 vocab_file/sp_model_file 互斥,推断 do_lower_case
  2. 按配置构建编码器:--bert_config_fileget_bert_encoder()--encoder_config_fileencoders.build_encoder()export_tfhub_lib.py);
  3. _create_model() 把编码器的命名输入转为 dict 输入(只有 dict 形式在 SavedModel 恢复后才可接受 dict 调用),前向一次得到输出 dict,并注入 default 别名,包成 tf_keras.Model(inputs=..., outputs=...)export_tfhub_lib.py);
  4. 构造双命名空间检查点(model=/encoder=)恢复权重并断言匹配;
  5. 视需要复制 NSP pooler dense 层;
  6. 把词表资产与 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
  • 发布前完成微调对拍验证。
登录后查看全文
热门项目推荐
相关项目推荐