首页
/ TensorFlow Models LaBSE:多语言句子嵌入的训练、配置与 TF-Hub 导出实战指南

TensorFlow Models LaBSE:多语言句子嵌入的训练、配置与 TF-Hub 导出实战指南

2026-09-04 09:13:06作者:何举烈Damon

本文基于 LaBSE 项目文档 展开,系统讲解 LaBSE(Language-agnostic BERT Sentence Embedding,语言无关的 BERT 句子嵌入)在 TensorFlow 官方 Model Garden 中的完整落地方案:从环境准备、数据格式、训练命令与 YAML 配置逐项解析,到双编码器任务、损失函数、Keras 模型的源码级实现,以及最终通过 export_tfhub.py 导出 SavedModel 的流程。读完后你可以独立完成 LaBSE 风格的双编码器句子嵌入训练、配置调参与模型导出。

一、LaBSE 是什么,以及本仓库实现的边界

LaBSE 是一种跨语言句子嵌入方法,核心思想是用 BERT 编码器把不同语言的句子映射到同一语义空间,从而支持跨语言检索、聚类和相似度计算。本仓库 official/projects/labse/ 目录下包含作者提供的官方实现与实验定义,但原文档明确了两点边界,使用时必须注意:

  1. 跨加速器全局 batch softmax 未实现:论文中在全局大 batch(跨 TPU/加速器)上计算的 softmax 对比学习在代码中未实现,因此该实现还不能完全复现论文config_labse.py 中的实验注册注释也重申了这一点:"this experiment does not use cross-accelerator global softmax so it does not reproduce the exact LABSE training");
  2. 训练数据不公开:受数据政策限制,作者无法发布 LaBSE 的预训练与微调数据,需要自行准备符合格式要求的多语言数据。

另外,若已训好模型并希望直接使用,可以在 TF Hub 公共模型库中查找 google/LaBSE(版本 2);该 Hub 上的 SavedModel 正是由本仓库的 export_tfhub.py 导出的。

二、环境要求与版本验证

根据文档,代码要求 TensorFlow 2.8.0(此后持续跟踪 TensorFlow 最新发行版)和 Python 3.7+。可按原文档给出的命令验证环境:

python --version
python -c 'import tensorflow as tf; print(tf.__version__)'

此外,运行模型时还需要把仓库的 models 目录加入 Python 路径,通用的运行方式可参考 official/README.md 中的说明。

三、目录结构与文件职责

official/projects/labse/ 目录结构简洁,每个文件职责清晰:

文件 作用
train.py 训练程序入口,注册 LaBSE 配置并启动训练
config_labse.py 注册 labse/train 实验,定义优化器/学习率默认值
experiments/labse_base.yaml 数据、模型与训练器的主配置
experiments/labse_bert_base.yaml BERT 编码器结构超参数
export_tfhub.py 将训练结果与预处理导出为 TF-Hub SavedModel
export_tfhub_test.py 验证导出模型与原模型权重、行为一致

训练入口 train.py 本体非常短:调用 tfm_flags.define_flags() 定义通用命令行参数后,app.run(train.main) 交给 NLP 训练框架 official/nlp/train.py 驱动整个实验流程。

四、数据准备:预训练与微调格式

预训练数据

预训练数据要求多语言,格式与 BERT 预训练数据相同(即 BERT pretraining 的 tensorflow.Example 格式,含掩码语言模型与下一句预测信号)。

微调数据

微调数据是成对句子的 tensorflow.Example,左右句分别放在 src_rawtgt_raw 两个 bytes_list 特征中。原文档给出的完整格式如下:

{   # (tensorflow.Example)
  features: {
    feature: {
      key  : "src_raw"
      value: {
        bytes_list: {
          value: [ "Foo. " ]
        }
      }
    }
    feature: {
      key  : "tgt_raw"
      value: {
        bytes_list: {
          value: [ "Bar. " ]
        }
      }
    }
  }
}

这一格式与配置文件中 left_text_fields: ['src_raw']right_text_fields: ['tgt_raw'] 的设定一一对应(见下文 labse_base.yaml 解析),数据加载由 dual_encoder_dataloader.py 完成。

五、启动训练:完整命令与参数说明

生成好预训练/微调数据后,按原文档的命令启动训练:

TPU=local
VOCAB=???
INIT_CHECKPOINT=???
PARAMS="task.train_data.input_path=/path/to/train/data"
PARAMS="${PARAMS},task.train_data.vocab_file=${VOCAB}"
PARAMS="${PARAMS},task.validation_data.input_path=/path/to/validation/data"
PARAMS="${PARAMS},task.validation_data.vocab_file=${VOCAB}"
PARAMS="${PARAMS},task.init_checkpoint=${INIT_CHECKPOINT}"
PARAMS="${PARAMS},runtime.distribution_strategy=tpu"

python3 train.py \
  --experiment=labse/train \
  --config_file=./experiments/labse_bert_base.yaml \
  --config_file=./experiments/labse_base.yaml \
  --params_override=${PARAMS} \
  --tpu=${TPU} \
  --model_dir=/folder/to/hold/logs/and/models/ \
  --mode=train_and_eval

参数说明:

  • --experiment=labse/train:指定实验名,对应 config_labse.py 中通过 @exp_factory.register_config_factory("labse/train") 注册的工厂函数;
  • --config_file:命令依次传入两个 YAML——labse_bert_base.yaml 定义 BERT 编码器结构,labse_base.yaml 定义模型训练策略、数据与训练器参数,二者共同组成最终实验配置;
  • --params_override:以 key=value 逗号分隔的形式覆盖配置项,用于注入训练/验证数据路径、词表文件、初始化 checkpoint 和分布式策略,避免把数据路径硬编码进 YAML;
  • --tpu:加速器地址,local 或 TPU 主机名/地址;
  • --model_dir:日志与模型保存目录;
  • --mode=train_and_eval:训练与评估同时执行。

注意:INIT_CHECKPOINT 应指向用 LaBSE 词表预训练好的 BERT checkpoint(配置中 init_checkpoint 的注释原文为 "the pre-trained BERT checkpoint using the labse vocab.")。加载逻辑在 DualEncoderTask.initialize 中:通过 pretrain2finetune_mapping = {'encoder': model.checkpoint_items['encoder']} 只映射 encoder 子网络,并调用 status.expect_partial().assert_existing_objects_matched() 做部分匹配恢复,因此即使外层双编码器结构在预训练中不存在也能正常加载。

六、配置文件逐项解析

labse_base.yaml:模型、数据与训练器

labse_base.yaml 是训练的核心配置,关键取值如下:

配置段 默认取值 含义
task.model bidirectional true 双向训练,左右句对都计算对比损失(对应 DualEncoder 同时输出 left_logits/right_logits
task.model max_sequence_length 32 句子最大 token 长度,短句嵌入任务用 32 而非 BERT 的 512
task.model logit_scale 100 点积 logits 的缩放系数,放大相似度差异
task.model logit_margin 0.3 正负样本对的附加 margin(additive margin 对比学习)
task.train_data global_batch_size 4096 全局 batch,in-batch 负采样的规模由此决定
task.train_data left_text_fields / right_text_fields ['src_raw'] / ['tgt_raw'] 与第四节数据格式对应
task.train_data seq_length 32 输入序列长度
task.train_data shuffle_buffer_size / cycle_length 1000 / 4 打乱缓冲区与并行预处理度
task.train_data lower_case false 不做小写化
task.validation_data global_batch_size 32000 验证 batch 更大
task.validation_data sharding true 验证数据启用分片读取
task.train_data / validation_data drop_remainder true / false 训练丢尾部、验证保留尾部
trainer optimizer_config.optimizer adamw AdamW,beta_1=0.9beta_2=0.999epsilon=1e-5gradient_clip_norm=100
trainer learning_rate.polynomial 初值 1e-4decay_steps=500000end=0.0power=1.0 多项式线性衰减
trainer warmup.polynomial warmup_steps=5000 学习率线性预热
trainer train_steps / steps_per_loop / checkpoint_interval / validation_interval 500000 / 1000 / 1000 / 1000 训练 50 万步,每 1000 步做一轮循环、存 checkpoint、评估一次(validation_steps: 100

labse_bert_base.yaml:编码器结构

labse_bert_base.yaml 仅覆盖编码器超参数,对应标准 BERT-base 结构:

task:
  model:
    encoder:
      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: 501153

值得注意的点是 vocab_size: 501153——远大于英文 BERT 的 30522,这正是 LaBSE 多语言词表的规模,也解释了为什么 INIT_CHECKPOINT 必须使用 LaBSE 词表预训练的 BERT,且微调数据要指定对应的 VOCAB 文件。

config_labse.py 中的默认值与覆盖

LaBSEOptimizationConfig 定义了实验级默认优化策略:AdamW(权重衰减)、多项式学习率(初始 1e-4decay_steps=1e6、末端 0)、1 万步预热。而 labse_train() 工厂函数在此基础上进一步覆盖为初始学习率 3e-5 并重置多项式预热,同时要求 task.train_data.is_trainingtask.validation_data.is_training 不得为 None(通过 restrictions 约束)。从源码结构看,YAML 与代码默认值共同构成最终配置,这也是命令行可以用 --params_override 精细覆盖任意键的基础。

七、核心实现:双编码器任务与模型

1. 任务层:in-batch 对比学习

LaBSE 训练复用 NLP 库中的双编码器任务 official/nlp/tasks/dual_encoder.py。其 build_model 支持两种编码器来源:hub_module_url(从 Hub 加载)或按配置构建本地编码器(二者只能指定其一),最终构造 models.DualEncoder(..., output='logits') 进入训练态。

损失函数 build_losses 是典型的 in-batch 对比学习:

  • tf.range(batch_size) 作为"正样本即对角线"的隐式标签;
  • 对左句 logits 计算 sparse_softmax_cross_entropy_with_logits
  • bidirectional=true 时(LaBSE 配置即为 true),对右句 logits 再算一次并相加,实现句子对的对称训练。

评估指标由 build_metrics 生成:默认 eval_top_k=(1, 3, 10),即同时统计 left_recall_at_{1,3,10}(双向时还有 right_recall_at_*),训练中即可直接观察跨语言检索的召回率。

2. 模型层:归一化、缩放与 margin

Keras 模型 DualEncoder 接收一个 transformer 编码器网络,构建左右两个塔:

  • normalize=True(LaBSE 训练与导出默认开启),对 pooled_outputtf.nn.l2_normalizeL77-L80),使嵌入落在单位球面上,点积退化为余弦相似度;
  • 训练态(output='logits')通过 MatMulWithMargin 层(official.nlp.modeling.layers 模块)计算左右句对点积矩阵,并应用 logit_scale=100 的缩放与 logit_margin=0.3 的正样本惩罚(margin 对比学习的实现细节见该层文档中引用的 additively margin 论文);
  • 推理态(output='predictions')只保留左塔,输入名沿用旧版 BERT Hub 模块的 input_word_ids/input_mask/input_type_ids,输出 sequence_outputpooled_output,保持与既有 BERT Hub 模型的调用习惯一致(L114-L120);
  • checkpoint_items 属性(L158-L161)把 encoder 暴露为可 checkpoint 项,这正是第六节 initialize 部分恢复预训练权重的接口。

八、导出 TF-Hub SavedModel

训练完成后,用 export_tfhub.py 分两步导出。文档头部的官方用法示例:

LaBSE_DIR=<Your LaBSE model dir>
# Step 1: 导出核心 LaBSE 模型
python3 ./export_tfhub.py \
  --bert_config_file ${LaBSE_DIR:?}/bert_config.json \
  --model_checkpoint_path ${LaBSE_DIR:?}/labse_model.ckpt \
  --vocab_file ${LaBSE_DIR:?}/vocab.txt \
  --export_type model --export_path /tmp/labse_model
# Step 2: 导出配套的预处理模块(务必使用相同的关键参数)
python3 ./export_tfhub.py \
  --vocab_file ${LaBSE_DIR:?}/vocab.txt \
  --export_type preprocessing --export_path /tmp/labse_preprocessing

主要参数(flags 定义):

参数 默认值 说明
--export_type model model(核心模型)或 preprocessing(预处理模块)
--export_path 必填 导出的 SavedModel 目标路径
--bert_config_file / --bert_tfhub_module 定义 BERT 核心层,二选一;后者设置时前者被忽略
--model_checkpoint_path 模型导出时必填的 checkpoint 路径
--vocab_file 词表文件,modelpreprocessing 两种导出都需要
--do_lower_case 自动推断 若为 None,则根据 vocab_file 文件名中是否含 uncased 自动决定小写化
--default_seq_length 128 预处理顶层 preprocess 方法与 bert_pack_inputs 子对象的默认序列长度
--tokenize_with_offsets False 是否额外导出 .tokenize_with_offsets 子对象
--normalize True 是否对嵌入(pooled_output)做归一化

实现上,mainmodel 分支调用 export_labse_model 恢复 encoder checkpoint 并保存(同时把 vocab 作为 tf.saved_model.Assetdo_lower_case 作为不可训练 tf.Variable 一并写入 SavedModel);preprocessing 分支复用 export_tfhub_lib.export_bert_preprocessing——注释明确说明 "LaBSE is still a BERT model, reuse the export_bert_preprocessing here"。

export_tfhub_test.py 提供了导出正确性的可验证依据:它用微型 BERT 配置构建模型并保存 checkpoint,导出后经 hub.KerasLayer 恢复,断言可训练权重逐一全等、pooled_output/sequence_output 形状正确,并验证 training=True 时 dropout 生效(20 次前向的标准差显著大于 1e-3)。这说明导出的 Hub 模型与源码模型在数值与训练行为上是一致的。

九、小结:适用前提与使用建议

  • 该实现不做跨加速器全局 softmax,训练时负样本只来自单个设备的 batch(配置中 global_batch_size=4096 即是这个规模),效果与论文原设定存在差距,复现实验时需知晓这一点;
  • 训练链路完全构建在 NLP 建模库之上:编码器、任务、数据加载分别位于 official/nlp/modeling/official/nlp/tasks/official/nlp/data/,若要更换编码器结构或数据字段,从 config_labse.pyDualEncoderConfig 与两个 YAML 入手即可;
  • 需要现成模型时,直接使用 TF Hub 上的 google/LaBSE(由本仓库导出脚本生成),或按本文第五、八节自行训练并导出。
登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.12 K
2.72 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
528
588
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
906
1.83 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
854
1.34 K
docsdocs
暂无描述
Markdown
891
5.79 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.53 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.34 K
1.45 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
988
506
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
540
384