使用 Sentence Transformers 訓練速度快 400 倍的靜態嵌入模型
TL;DR
Hugging Face 發布了一套訓練靜態嵌入模型的配方,這些模型在 CPU 上的運行速度比標準基於 transformer 的嵌入模型快 100x–400x,同時保留了至少 85% 的品質。此外,還發布了兩個具體的模型(用於英文檢索的 static‑retrieval‑mrl‑en‑v1 和用於多語言相似度的 static‑similarity‑mrl‑multilingual‑v1),並附帶訓練腳本、評估結果和 Weights & Biases 日誌。
方法
該方法將對比學習 (contrastive learning) 與 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 進行零樣本 (zero-shot) 檢索評估;多語言模型依賴 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 僅會降低 1.47% 的 NDCG@10 (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 加速,同時保留常見基於 transformer 嵌入模型至少 85% 的品質。發布的模型能夠以極小的準確度損失,實現高效的裝置端、瀏覽器端和邊緣運算應用場景。
下一步
使用者可以將現有的 Sentence Transformer 模型替換為 static‑retrieval‑mrl‑en‑v1 或 static‑similarity‑mrl‑multilingual‑v1,或者在特定任務的數據上訓練自己的靜態嵌入。潛在的改進包括:硬負樣本挖掘 (hard-negative mining)、模型集成 (model souping)、課程學習 (curriculum learning)、引導式 batch 內負樣本過濾、種子優化的隨機初始化、tokenizer 重新訓練、透過 CachedMultipleNegativesRankingLoss 進行梯度緩存,以及從更大的編碼器進行模型蒸餾。