使用 Sentence Transformers 训练速度提升 400 倍的静态嵌入模型

TL;DR

Hugging Face 发布了一种训练静态嵌入模型的方法,该模型在 CPU 上的运行速度比标准基于 transformer 的嵌入快 100x–400x,同时保留了至少 85% 的质量。此外,还发布了两个具体的模型(用于英语检索的 static‑retrieval‑mrl‑en‑v1 和用于多语言相似度的 static‑similarity‑mrl‑multilingual‑v1),以及训练脚本、评估结果和 Weights & Biases 日志。

方法

该方法将对比学习与 MultipleNegativesRankingLoss 以及可选的 Matryoshka Representation Learning 相结合,用于训练静态编码器。这些编码器执行简单的 token‑embedding 查找,而不是基于 attention 的编码,从而在损失极小质量的情况下实现了数个数量级的加速。

训练详情

要求

训练使用 Sentence Transformers 库,包含以下组件:dataset、loss function、training arguments、evaluator 和 trainer。

模型灵感

目标是两个模型:一个仅限英语的检索模型和一个多语言通用相似度模型,两者都使用 StaticEmbedding 模块进行初始化,该模块封装了一个基于 BERT 的 tokenizer(bert‑base‑uncased 或 bert‑base‑multilingual‑uncased),嵌入维度为 1024。

数据集选择

对于英语检索,选择了 30 个数据集,包括 gooaq、msmarco (triplet)、squad、s2orc (title‑abstract‑pair)、allnli (triplet)、paq、trivia_qa、msmarco_10m、swim_ir (en)、pubmedqa (triplet‑20)、miracl (en‑triplet‑all)、mldr (en‑triplet‑all) 和 mr_tydi (en‑triplet‑all)。对于多语言相似度,结合了具有平行句子的数据集(wikititles、tatoeba、talks、europarl、global_voices、jw300、muse、wikimatrix、opensubtitles)和正样本对数据集(stackexchange‑duplicates、quora‑duplicates、wikianswers‑duplicates、allnli、simple_wiki、altlex、flickr30k_captions、coco_captions、nli_for_simcse、negation),总计 30 个训练数据集。

损失函数

选择 MultipleNegativesRankingLoss 是因为 2048 的 batch size 可以适配 RTX 3090,从而避免了 CachedMultipleNegativesRankingLoss 的开销以及 GISTEmbedLoss 中 guide model 带来的减速。在此基础上应用了 MatryoshkaLoss,维度为 [32, 64, 128, 256, 512, 1024]。

训练参数

两个模型均使用:num_train_epochs=1, per_device_train_batch_size=2048, per_device_eval_batch_size=2048, learning_rate=2e-1, warmup_ratio=0.1, bf16=True, batch_sampler=BatchSamplers.NO_DUPLICATES, multi_dataset_batch_sampler=MultiDatasetBatchSamplers.PROPORTIONAL, eval_strategy=steps, eval_steps=250 (英语) 或 1000 (多语言), save_strategy=steps, save_steps 与 eval_steps 一致, save_total_limit=2, logging_steps 与 eval_steps 一致, logging_first_step=True,以及一个特定运行的输出目录。

评估器

英语检索模型使用 NanoBEIREvaluator 进行零样本检索评估;多语言模型依靠 MTEB 任务(STS、Classification、Pair Classification)进行评估。

硬件

训练是在 RTX 3090 GPU、i7‑13700K CPU 和 32 GB RAM 上进行的。

整体训练脚本

提供的脚本负责加载数据集、实例化 StaticEmbedding 模型、结合 MatryoshkaLoss 组合损失函数、设置训练参数、可选地运行评估器、使用 SentenceTransformerTrainer 进行训练并保存最终模型。英语检索脚本耗时 17.8 小时,消耗 2.6 kWh,排放 1 kg CO₂;多语言脚本耗时 3.1 小时,消耗 0.5 kWh,排放 0.2 kg CO₂。

使用方法

两个模型都可以通过 SentenceTransformer 使用模型名称和 device="cpu" 来加载。推理过程与标准的 Sentence Transformers 相同:model.encode 返回嵌入,model.similarity 计算余弦相似度。truncate_dim 参数支持 Matryoshka 式的降维(例如 truncate_dim=256)。这些模型可以开箱即用地与 LangChain、LlamaIndex、Haystack 和 txtai 配合使用。

性能

英语检索

在 NanoBEIR 上,static‑retrieval‑mrl‑en‑v1 的 NDCG@10 达到 0.5032,是 all‑mpnet‑base‑v2 (0.5757) 分数的 87.4%。在 CPU 上,它每秒处理 107,419.51 个句子,比 all‑mpnet‑base‑v2 (270.40 句子/秒) 快 397 倍。在 GPU 上,它每秒处理 97,171.47 个句子,比 all‑mpnet‑base‑v2 (4043.13 句子/秒) 快 24 倍。Matryoshka 评估显示,将维度减半至 512 仅会使 NDCG@10 下降 1.47% (0.5032 → 0.4957)。

多语言相似度

相对于 multilingual‑e5‑small,static‑similarity‑mrl‑multilingual‑v1 在 STS 上的得分为 92.3%,在 Pair Classification 上为 95.52%,在 Classification 上为 86.52%。它在 CPU 上的速度比 multilingual‑e5‑small 快约 125 倍,在 GPU 上快 10 倍。Matryoshka 评估表明,将维度降低到 256(缩小 4 倍)仅会导致英语 STS 性能下降 0.56%。

结论

使用本文介绍的配方训练的静态嵌入模型可实现 100x–400x 的 CPU 加速和 10x–25x 的 GPU 加速,同时保留至少 85% 的常用基于 transformer 嵌入的质量。发布的模型能够在极小的精度损失下,实现高效的设备端、浏览器端和边缘计算用例。

下一步工作

用户可以将现有的 Sentence Transformer 模型替换为 static‑retrieval‑mrl‑en‑v1 或 static‑similarity‑mrl‑multilingual‑v1,或者在特定任务数据上训练自己的静态嵌入。潜在的改进包括:难负样本挖掘 (hard-negative mining)、模型集成 (model souping)、课程学习 (curriculum learning)、引导式 batch 内假负样本过滤、种子优化随机初始化、tokenizer 重训练、通过 CachedMultipleNegativesRankingLoss 进行梯度缓存,以及从更大的编码器进行模型蒸馏。

Sources