首页
/ tensorflow/models 中的 MobileBERT:紧凑 BERT 的 TF2 实现、渐进蒸馏与预训练模型加载全解

tensorflow/models 中的 MobileBERT:紧凑 BERT 的 TF2 实现、渐进蒸馏与预训练模型加载全解

2026-09-04 12:54:22作者:乔或婵

本篇指南围绕 tensorflow/models 仓库中 official/projects/mobilebert 目录的 MobileBERT 项目展开:先讲清 MobileBERT 这一"薄版 BERT_LARGE"的网络结构与 TF2/Keras 实现,再完整覆盖预训练模型的规格与加载方式,并结合仓库源码深入解析其渐进式知识蒸馏训练流水线、TF1 到 TF2 的检查点转换工具以及 TF-Hub 模型导出流程。读完后,你将能够基于本仓库代码加载 MobileBERT 编码器、复现其蒸馏训练配置,并理解每个配置项在源码中的实际作用。

一、MobileBERT 是什么

MobileBERT 是一个 BERT_LARGE 的"精简版本(thin version)",它在网络中引入瓶颈(bottleneck)结构,并精心平衡了自注意力(self-attention)与前馈网络(feed-forward networks)之间的容量配比。按照 项目 README 的说明,其训练方式是两阶段的:

  1. 先训练一个特殊设计的教师模型(teacher)——在 BERT_LARGE 中嵌入倒瓶颈(inverted-bottleneck)结构的变体;
  2. 再将该教师模型的"知识"通过蒸馏(knowledge transfer)传递给 MobileBERT 学生模型(student)。

官方实证研究表明,MobileBERT 相比 BERT_BASE 参数量缩小 4.3 倍、推理速度提升 5.5 倍,同时在主流基准上取得了有竞争力的结果。本仓库包含 MobileBERT 的 TensorFlow 2.x 实现,核心代码分布在 NLP modeling 库与 official/projects/mobilebert 项目目录两处。

二、TF2 网络实现:编码器与基础层

README 明确给出了两个核心实现文件的位置(均为 TF2 tf.keras API 重新实现):

2.1 MobileBERTEncoder:函数式 Keras 编码器

MobileBERTEncoder 是一个 tf_keras.Model 子类,采用 Keras functional API 构图。其关键构造参数与默认值如下(源自构造函数签名,L26-L46):

参数 默认值 含义
word_vocab_size 30522 词表大小
word_embed_size 128 词嵌入维度(注意远小于 hidden_size)
type_vocab_size 2 句子类型数
max_sequence_length 512 最大输入序列长度
num_blocks 24 Transformer 块数量
hidden_size 512 隐藏层宽度
num_attention_heads 4 注意力头数
intermediate_size 512 FFN 中间层宽度
intra_bottleneck_size 128 瓶颈宽度
key_query_shared_bottleneck True 是否共享 K/Q 的线性变换
num_feedforward_networks 4 每个块堆叠的 FFN 数量
normalization_type no_norm 归一化类型,仅支持 no_normlayer_norm
classifier_activation False [CLS] 池化是否加 tanh 激活

编码器的前向结构在 L131-L164 中完成:先由 MobileBertEmbedding 生成嵌入输出,随后顺序堆叠 num_blocksMobileBertTransformer 层,最后取第一个 token 作为 pooled_output。整个编码器输出一个字典,包含 sequence_output(完整序列表示)、pooled_output([CLS] 表示)、encoder_outputs(各层输出)与 attention_scores(各层注意力分数)——后两者正是蒸馏训练中做逐层对齐的关键。

源码中对 normalization_type 的注释值得注意:no_norm 表示学生模型使用的逐元素线性变换(源自原 MobileBERT 论文的建议),layer_norm 则用于教师模型。也就是说,同一份实现同时承载了"学生/教师"两种网络变体。此外,input_mask_dtype 参数默认 int32;若下游要做 TF Lite 量化(不支持 Cast op),可将其设为 float32 以规避计算图中的类型转换。

2.2 MobileBertEmbedding:三语嵌入 + 局部卷积式输入

MobileBertEmbedding 包含词嵌入、句子类型嵌入与位置嵌入三部分,并在 call 方法 中实现了一个 MobileBERT 论文的特色设计:局部卷积输入(trigram input)——将词嵌入沿序列维度向左右各外扩一个 token 并沿特征维拼接(相当于宽度为 3 的局部卷积),再通过 embedding_projection 投影到 output_embed_size,最后叠加位置嵌入与类型嵌入,经归一化与 dropout 输出。这一设计让每个位置的表示显式包含相邻词信息,为模型在浅层隐式获得了局部上下文能力。

2.3 MobileBertTransformer:瓶颈化的 Transformer 块

MobileBertTransformer 实现了一个带瓶颈与倒瓶颈结构的 Transformer 块,其内部按顺序组织为五个子模块(self.block_layers,见 L244-L325):

  1. input bottleneck(bottleneck_input:用 EinsumDensehidden_size 压到 intra_bottleneck_size,再过一个归一化层;
  2. K/Q 共享瓶颈(kq_shared_bottleneck:当 key_query_shared_bottleneck=True 时,Key 与 Query 共用同一个投影,Value 仍走原始张量,从而省掉一组线性层参数(见 call 方法 L380-L395);
  3. attention:标准 MultiHeadAttention,头维度为 intra_bottleneck_size / num_attention_heads
  4. stacked FFN(ffnnum_feedforward_networks 个堆叠的前馈网络,每个为"intermediate_dense(带激活)→ output_dense → 归一化";
  5. output bottleneck(bottleneck_output:把张量从瓶颈宽度恢复回 hidden_size,接 dropout 与归一化。

构造函数还包含一个硬性约束(L238-L242):intra_bottleneck_size 必须是 num_attention_heads 的整数倍,否则无法均分到每个注意力头。

归一化层由辅助函数 _get_norm_layer 根据 normalization_type 返回:layer_norm 对应标准 LayerNormalization,而 no_norm 对应 NoNorm——一个只学习逐元素缩放 gamma 与偏置 beta 的轻量层,这正是学生模型省掉完整归一化计算、换取端侧推理速度的实现细节。

三、预训练模型一览

README 将原 TF 1.x 预训练英语 MobileBERT 检查点转换为了 TF 2.x 检查点(与上述实现兼容),并额外提供了用多语言 Wiki 数据训练的多语言 MobileBERT 检查点;两者均导出了 TF-Hub SavedModel。官方给出的模型规格表如下:

模型 配置 参数量 训练数据 指标
MobileBERT uncased English uncased_L-24_H-128_B-512_A-4_F-4_OPT 25.3 Million Wiki + Books Squad v1.1 F1 90.0,GLUE 77.7
MobileBERT cased Multi-lingual multi_cased_L-24_H-128_B-512_A-4_F-4_OPT 36 Million Wiki XNLI (zero-shot): 64.7

配置命名中的字母对应论文中的超参数记号(L-24 层、H-128 瓶颈宽度、B-512 隐藏宽度、A-4 注意力头、F-4 FFN 个数),与下文学生模型 yaml 中的字段一一对应。TF-Hub 模型可分别按 tensorflow/mobilebert_en_uncased_L-24_H-128_B-512_A-4_F-4_OPT/1(英语)与 tensorflow/mobilebert_multi_cased_L-24_H-128_B-512_A-4_F-4_OPT/1(多语言)两个模型名检索使用。

四、从检查点恢复 MobileBERT

README 给出的官方加载示例(可复制使用):

import tensorflow as tf
from official.nlp.projects.mobilebert import model_utils

bert_config_file = ...
model_checkpoint_path = ...

bert_config = model_utils.BertConfig.from_json_file(bert_config_file)

# `pretrainer` is an instance of `nlp.modeling.models.BertPretrainerV2`.
pretrainer = model_utils.create_mobilebert_pretrainer(bert_config)
checkpoint = tf.train.Checkpoint(**pretrainer.checkpoint_items)
checkpoint.restore(model_checkpoint_path).assert_existing_objects_matched()

# `mobilebert_encoder` is an instance of
# `nlp.modeling.networks.MobileBERTEncoder`.
mobilebert_encoder = pretrainer.encoder_network

这段代码背后的实现值得拆开看:

  • BertConfig 是一个独立的轻量配置类(不依赖 NLP configs 体系),除了常规 BERT 参数外,还包含 MobileBERT 专属字段:trigram_inputuse_bottleneckintra_bottleneck_sizeuse_bottleneck_attentionkey_query_shared_bottlenecknum_feedforward_networksnormalization_typeclassifier_activationfrom_dict 中有两个自动补全逻辑:embedding_size 缺省时取 hidden_sizeintra_bottleneck_size 缺省时也取 hidden_sizeL114-L117)。
  • create_mobilebert_pretrainer 负责把 config 映射为 MobileBERTEncoder + MobileBertMaskedLM(共享词嵌入表),再包装进 BertPretrainerV2,并调用一次前向以强制创建全部变量——这一步保证随后的 checkpoint.restore(...).assert_existing_objects_matched() 能逐对象匹配校验。

五、渐进蒸馏训练流水线(源码级解析)

README 提到蒸馏训练,而仓库提供了完整可运行的流水线,入口是 run_distillation.py,核心逻辑在 distillation.py

5.1 三阶段配置体系

蒸馏行为由三组 dataclass 配置描述:

  • LayerWiseDistillConfig:逐层蒸馏阶段。默认 num_steps=10000initial_learning_rate=1.5e-3hidden_distill_factor=100.0beta_distill_factor=5000.0gamma_distill_factor=5.0if_transfer_attention=Trueattention_distill_factor=1.0;其中 transfer_teacher_layers 允许把层数更多的教师映射到学生(例如把 24 层教师压到 6 层学生时设为 [3, 7, 11, 15, 19, 23]),为 None 时要求师生层数相同;
  • PretrainDistillConfig:最后的预训练对齐阶段,默认 num_steps=500000warmup_steps=10000、学习率从 1.5e-3 衰减到 1.5e-7if_use_nsp_loss=Truedistill_ground_truth_ratio=0.5
  • BertDistillationProgressiveConfig:继承 ProgressiveConfig,含 if_copy_embeddings(是否把教师词嵌入直接拷贝给学生)及上述两个子配置。

任务级配置 BertDistillationTaskConfig 中,教师与学生默认都是 encoders.EncoderConfig(type='mobilebert')PretrainerConfig,另含教师初始检查点路径 teacher_model_init_checkpoint 与训练/验证数据配置。

5.2 逐层对齐的损失函数

BertDistillationTask 继承 ProgressivePolicy,其阶段数等于学生层数加 1(num_stages):前 N 个阶段每个阶段只训练学生的第 N 层,最后 1 个阶段做完整预训练对齐。

  • 前 N 阶段:build_modelbuild_sub_encoderL106-L123)分别切出"教师第 K 层为止"与"学生第 K 层为止"的子编码器,模型输出四路特征;
  • 损失(build_losses)由四部分构成:
    • 特征迁移损失:对师生隐藏态各过一个不可训练的 LayerNormalization 后计算 MSE,乘以 hidden_distill_factor(默认 100);
    • β/γ 分布损失:分别对齐师生特征的均值平方差(beta_distill_factor 默认 5000)与方差绝对差(gamma_distill_factor 默认 5);
    • 注意力迁移损失:教师注意力 softmax 与 log_softmax(学生注意力) 的 KL 散度,乘以 attention_distill_factor
    • 总损失再除以 stage_id + 1 做阶段化缩放,避免深阶段损失量级偏大。
  • 最后阶段:学生整体对教师的 MLM 输出做软标签蒸馏。build_losses L421-L449 中,真实 one-hot 标签与教师 MLM 的 softmax 标签按 distill_ground_truth_ratio(默认 0.5)线性混合:lm_label = gt_ratio * lm_label + (1-gt_ratio) * teacher_labels,学生以交叉熵拟合该混合标签;若数据含 next_sentence_labels 则叠加 NSP 分类损失。

训练开始时,initialize 方法(L581-L605)从 teacher_model_init_checkpoint 加载教师权重(支持传目录,自动取最新检查点),并把教师嵌入层权重直接拷贝给学生。

5.3 训练入口与优化器

run_distillation.py 定义了默认优化器:LAMB(weight_decay_rate=0.01,排除 LayerNorm/bias/normclipnorm=1.0)+ 多项式学习率衰减 + 线性 warmup;get_exp_config 中默认 train_steps=740000checkpoint_interval=20000main 中支持 gin 参数、混合精度策略与 TPU/GPU 分布式策略。参数覆盖逻辑(config_override)支持 --config_file(一级覆盖)与 --params_override(二级覆盖),最后 validate + lock 并打印最终参数。

5.4 官方实验 yaml:教师/学生配置对比

experiments/ 目录下的三份 yaml 是完整可参考的配置模板:

  • en_uncased_teacher.yaml:教师——intermediate_size: 4096intra_bottleneck_size: 1024hidden_activation: gelunum_feedforward_networks: 1normalization_type: layer_normkey_query_shared_bottleneck: falsehidden_dropout_prob: 0.1
  • en_uncased_student.yaml:学生——intermediate_size: 512intra_bottleneck_size: 128hidden_activation: relunum_feedforward_networks: 4normalization_type: no_normkey_query_shared_bottleneck: truehidden_dropout_prob: 0.0
  • mobilebert_distillation_en_uncased.yaml:把上述师生架构合并,并给出数据与训练参数:global_batch_size: 2048seq_length: 512max_predictions_per_seq: 20use_next_sentence_label: trueuse_position_id: false(学生额外配了 next_sentence 分类头:inner_dim: 512num_classes: 2、tanh 激活);蒸馏侧 layer_wise_distill_config.num_steps: 10000pretrain_distill_config.num_steps: 500000train_steps: 740000max_to_keep: 10

从两份模型 yaml 的对比可以直观看到"倒瓶颈教师 + 薄学生"的设计:教师把宽度花在单个大 FFN(4096)与宽瓶颈(1024)上,学生则用 4 个小 FFN(512)与窄瓶颈(128)换取更低的端侧算力需求,这与 README 描述的"自注意力与前馈网络之间的平衡"完全一致。

六、TF1 到 TF2 检查点转换工具

README 提到官方将原 TF 1.x 英语检查点转换为了 TF 2.x 检查点,仓库中的实现是 tf2_model_checkpoint_converter.py。该脚本的命令行参数(L29-L40):

  • --bert_config_file:定义核心 MobileBERT 层的 JSON 配置文件;
  • --tf1_checkpoint_path:TF1 检查点路径;
  • --tf2_checkpoint_path:输出 TF2 检查点路径;
  • --use_model_prefix:当转换后的检查点用于子类化(subclass)实现的模型时打开,用模型名作为变量前缀。

转换流程为:先用 _NAME_REPLACEMENT 规则表做变量名迁移(如 bert/mobile_bert_encoder/embeddings/word_embeddingsmobile_bert_embedding/word_embedding/embeddingsattention/selfattention,以及按 num_feedforward_networks 动态展开的 LAST_FFN_LAYER_ID 占位替换,见 L197-L257);再按模式匹配对 attention 的 query/key/value kernel/bias 做按头 reshape(get_new_shape);排除 cls/seq_relationshipglobal_step;最后经 model_utils.create_mobilebert_pretrainer 重建 TF2 模型并调用 load_weights(...).assert_existing_objects_matched() 做名字级流式恢复,落盘为 V2 检查点(create_v2_checkpoint)。这套"改名 + 按头整形 + 断言匹配"的流程,也解释了为什么 --use_model_prefix 必须与使用侧的变量命名约定保持一致。

七、导出 TF-Hub SavedModel

官方把检查点导出为 TF-Hub SavedModel 的脚本是 export_tfhub.py。其命令行参数为 --bert_config_file--model_checkpoint_path--export_path--vocab_file--do_lower_case(默认 True)。导出逻辑(L36-L74):

  1. create_mobilebert_pretrainer 重建模型,并取 pretrainer.encoder_network
  2. 把编码器输出字典中 pooled_output 别名为 default,以兼容其他文本表示模型的调用习惯;
  3. 将 MLM 模型挂到 core model 的 mlm 属性上,并临时关闭 _auto_track_sub_layers,避免 MLM 权重被算进核心模型(源码注释标明是规避一个 TF bug 的临时做法);
  4. checkpoint.restore(model_checkpoint_path).assert_existing_objects_matched() 恢复权重后,把词表文件以 tf.saved_model.Asset 形式、do_lower_case 以不可训练 tf.Variable 形式打包进 SavedModel,save(format="tf") 完成 TF-Hub 格式导出。

这与 README 中"导出的 TF-Hub 模型可直接用于推理"的说明对应:导出的 SavedModel 同时携带了编码器核心输出、MLM 头与分词所需词表资产。

八、小结与延伸阅读

official/projects/mobilebert 目录构成了一个自洽的 MobileBERT 工程闭环:

如果你要在端侧场景使用轻量 BERT,建议的实操路径是:先用 README 第四节的恢复代码验证本地检查点与 MobileBERTEncoder 的变量匹配,再参考 experiments/ 下的 yaml 修改蒸馏配置做自定义压缩(例如借助 transfer_teacher_layers 压缩层数),最后用转换与导出脚本产出可分发的 TF2 检查点或 TF-Hub SavedModel。相关单元测试见 distillation_test.py,可用于验证你修改后的蒸馏逻辑。

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