Sentence Transformers 6.0 MultiVectorEncoder:訓練與微調指南

簡要重點

Hugging Face 在 Sentence Transformers v6.0 中新增了 MultiVectorEncoder 模型類型,並釋出逐步訓練程式碼,可在單一 RTX 3090 上於 14.5 小時內微調一個類 ColBERT 的晚期互動檢索器(例如 multi-vector-encoder/mLateOn-medical),達成 0.9139 NDCG@10 的成績——在醫療檢索基準上,超越 50 多種密集、稀疏、詞法與多向量基線模型的表現。


什麼是多向量(晚期互動)模型?

多向量模型會為每個詞元保留一個嵌入向量,而非將整個文字壓縮成單一向量。

  • 檢索使用 MaxSim 運算子:每個查詢詞元尋找最匹配的文件詞元,並對分數求和。
  • 詞元層級的匹配保留了細粒度的相關性訊號,這些訊號在單一向量模型中會被平均化,通常能提升檢索效果,但需付出更大的索引空間代價。
  • 附帶文章《使用 Sentence Transformers 的多向量(晚期互動)嵌入模型》詳細說明了架構、編碼、評分與索引方式。

"一個密集嵌入模型會將整個文字壓縮成單一向量,相似度僅是兩個摘要之間的一次點積。多向量模型……則為每個詞元保留一個小型向量,並使用 MaxSim 運算子將查詢與文件進行評分。" – 部分部落格內容


為何要微調多向量模型?

微調可讓模型適應特定領域(例如醫療、法律、程式碼)的詞彙、查詢風格與相關性概念。主要觀察如下:

  • 領域訊號 以詞元為單位捕捉,因此少量領域內資料即可帶來顯著提升。
  • 多數釋出的檢查點會將文件截斷至 180–512 個詞元;在長段落(MIRIAD 醫療資料集平均 941 個詞元)上,截斷可能導致最高 0.24 NDCG@10 的損失。
  • 非監督 檢查點(預對比但尚未監督)開始,其在領域適應上的表現持續優於完全監督的檢查點。
  • 整個微調流程可在單一消費級 GPU 上執行,使領域特定檢索器對大多數團隊都可取得。

訓練元件概覽

元件 功能
模型 MultiVectorEncoder 實例——可為預訓練檢查點或基於基礎 Transformer 建立的新模型
資料集 datasets.DatasetDatasetDict,包含查詢-文件對(或損失所需的其他格式)
損失 批次內負樣本的對比損失(CachedMultiVectorMultipleNegativesRankingLoss)或知識蒸餾損失
訓練參數 MultiVectorEncoderTrainingArguments——批次大小、學習率、提示等
評估器 MultiVectorInformationRetrievalEvaluator,用於 NDCG@10、acc@1 等指標
訓練器 MultiVectorEncoderTrainer 協調上述元件

以下各節可獨立閱讀;程式碼片段完整且可執行,只需執行 pip install -U "sentence-transformers[train]" 即可。


模型選擇與設定

微調現有檢查點

from sentence_transformers import MultiVectorEncoder
model = MultiVectorEncoder(
    "lightonai/mLateOn-unsupervised",
    model_kwargs={"torch_dtype": "float32"},
    processor_kwargs={"model_max_length": 8192},  # 允許完整長度文件
)
# 移除內建長度限制(例如 180–512 個詞元)
model[0].query_length = None
model[0].document_length = None
# 選用:跳過標點符號詞元以縮小索引
import string
model[2].skiplist_words = list(string.punctuation)
model[2].resolve_with_tokenizer(model.tokenizer)
  • 此檢查點已包含查詢/文件標記詞元、投影頭與評分跳過清單。
  • 提升限制後,模型可處理完整的 1,400 個詞元醫療段落。

從基礎 Transformer 建立新模型

from sentence_transformers import MultiVectorEncoder
model = MultiVectorEncoder(
    "answerdotai/ModernBERT-base",
    model_kwargs={"torch_dtype": "float32"},
)
  • 流程會自動加入隨機詞元級投影(128 維)與常見的 ColBERT 模組。
  • 即使強大的密集主幹模型(例如 Alibaba-NLP/gte-modernbert-base)在訓練 25 k 對後,也能達到接近檢查點的表現。

應選擇哪種起始點?

在 25 k MIRIAD 對上的實證比較:

檢查點 零樣本 NDCG@10 微調後 Δ
lightonai/mLateOn-unsupervised 0.9087 0.9398 +0.0311
lightonai/mLateOn 0.9277 0.9319 +0.0042
lightonai/LateOn-unsupervised 0.9026 0.9206 +0.0180
lightonai/LateOn 0.9185 0.9105 –0.0080
lightonai/GTE-ModernColBERT-v1 0.9198 0.9007 –0.0191
新增頭部於 gte-modernbert-base 0.9177
結論: 未經監督的檢查點適應性最佳;完全監督的檢查點常會退化。

資料集準備

訓練器接受任何 datasets.Dataset(Hub 或本機)。所需格式取決於損失函數:

  • 標籤欄位labelscore),若損失需要監督訊號。
  • 輸入欄位 依序排列;第一欄視為查詢,後續欄視為文件(除非透過 router_mapping 覆蓋)。

從 Hub 加載(MIRIAD 範例)

from datasets import load_dataset
train_dataset = load_dataset("tomaarsen/miriad-4.4M-split", split="train")
print(train_dataset)
# Dataset({features: ['question', 'passage_text'], num_rows: 4_467_542})
  • 每列提供一個 (question, passage_text) 對;段落平均長度為 941 個詞元。

本機 CSV/JSON 範例

from datasets import load_dataset
dataset = load_dataset("csv", data_files="my_file.csv")
# 或 JSON
# dataset = load_dataset("json", data_files="my_file.json")
  • 自訂預處理時,請使用 Dataset.from_dict

損失函數選擇

針對問答對,建議使用 批次內負樣本 損失搭配 GradCache:

from sentence_transformers.multi_vector_encoder.losses import CachedMultiVectorMultipleNegativesRankingLoss
loss = CachedMultiVectorMultipleNegativesRankingLoss(
    model=model,
    mini_batch_size=16,  # 用於記憶體的區塊大小;有效批次為 128(如部落格所示)
)
  • mini_batch_size 控制記憶體使用;有效對比批次大小維持在 128。
  • 縮放參數:請保留預設 scale=1.0 以適用 MaxSim 分數;使用密集模型的預設值(scale=20.0)將導致梯度飽和。
  • 知識蒸餾請參閱 MultiVectorDistillKLDivLoss(本部落格範例未使用)。

訓練參數

關鍵設定,可產出最佳結果:

from sentence_transformers import MultiVectorEncoderTrainingArguments
from sentence_transformers.base.sampler import BatchSamplers
args = MultiVectorEncoderTrainingArguments(
    output_dir="models/mLateOn-medical",
    num_train_epochs=1,
    per_device_train_batch_size=128,   # 透過 GradCache 的有效批次
    per_device_eval_batch_size=16,
    learning_rate=1e-4,
    warmup_steps=0.05,
    prompts={"question": "[Q] ", "passage_text": "[D] "},
    fp16=False,
    bf16=True,
    batch_sampler=BatchSamplers.NO_DUPLICATES,
    eval_strategy="steps",
    eval_steps=0.1,
    save_strategy="steps",
    save_steps=0.05,
    logging_steps=0.01,
    run_name="mLateOn-medical",
)
  • 提示 必須明確提供;模型不會自動套用儲存的標記。
  • 不設 max_length 可確保訓練時看到完整文件長度。
  • 在 5e‑6 → 2e‑4 的掃描中,較高的學習率(1e-4)表現最佳。

檢索評估器

最具資訊量的評估器是 MultiVectorInformationRetrievalEvaluator,其基於保留的查詢集與包含干擾項的語料庫建立:

from sentence_transformers.multi_vector_encoder.evaluation import MultiVectorInformationRetrievalEvaluator
# 建立語料庫、查詢與相關性映射(詳見部落格完整迴圈)
evaluator = MultiVectorInformationRetrievalEvaluator(
    queries=queries,
    corpus=corpus,
    relevant_docs=relevant_docs,
    name="miriad-dev",
    batch_size=16,
)
  • 包含困難的干擾段落以避免飽和;部落格將約 190 k 隨機訓練段落加入 10 k 黃金集。

完整訓練程式碼

以下程式碼可重現 mLateOn‑medical 模型:

import logging, string, traceback
from datasets import load_dataset
from sentence_transformers import (
    MultiVectorEncoder,
    MultiVectorEncoderModelCardData,
    MultiVectorEncoderTrainer,
    MultiVectorEncoderTrainingArguments,
)
from sentence_transformers.base.sampler import BatchSamplers
from sentence_transformers.multi_vector_encoder.losses import CachedMultiVectorMultipleNegativesRankingLoss
from sentence_transformers.multi_vector_encoder.evaluation import MultiVectorInformationRetrievalEvaluator

logging.basicConfig(format="%(asctime)s - %(message)s", level=logging.INFO)

def main():
    # 1️⃣ 加載非監督檢查點並解除限制
    model = MultiVectorEncoder(
        "lightonai/mLateOn-unsupervised",
        model_kwargs={"torch_dtype": "float32"},
        processor_kwargs={"model_max_length": 8192},
        model_card_data=MultiVectorEncoderModelCardData(
            language="en",
            license="apache-2.0",
            model_name="mLateOn 在 MIRIAD 醫療檢索上微調",
        ),
    )
    model[0].query_length = None
    model[0].document_length = None
    model[2].skiplist_words = list(string.punctuation)
    model[2].resolve_with_tokenizer(model.tokenizer)

    # 2️⃣ 加載 1 M 醫療問答對
    train_dataset = load_dataset("tomaarsen/miriad-4.4M-split", split="train").select(range(1_000_000))

    # 3️⃣ 定義損失(GradCache 批次內負樣本)
    loss = CachedMultiVectorMultipleNegativesRankingLoss(model=model, mini_batch_size=16)

    # 4️⃣ 輕量開發評估器(500 個查詢)
    eval_split = load_dataset("tomaarsen/miriad-4.4M-split", split="eval")
    corpus, queries, relevant_docs, passage_to_id = {}, {}, {}, {}
    for idx, row in enumerate(eval_split):
        if row["passage_text"] not in passage_to_id:
            pid = f"p{len(passage_to_id)}"
            passage_to_id[row["passage_text"]] = pid
            corpus[pid] = row["passage_text"]
        if idx < 500:
            qid = f"q{idx}"
            queries[qid] = row["question"]
            relevant_docs[qid] = {passage_to_id[row["passage_text"]]}  
    dev_evaluator = MultiVectorInformationRetrievalEvaluator(
        queries=queries, corpus=corpus, relevant_docs=relevant_docs, name="miriad-dev", batch_size=16
    )

    # 5️⃣ 訓練參數(參見前一節)
    args = MultiVectorEncoderTrainingArguments(
        output_dir="models/mLateOn-medical",
        num_train_epochs=1,
        per_device_train_batch_size=128,
        per_device_eval_batch_size=16,
        learning_rate=1e-4,
        warmup_steps=0.05,
        prompts={"question": "[Q] ", "passage_text": "[D] "},
        fp16=False,
        bf16=True,
        batch_sampler=BatchSamplers.NO_DUPLICATES,
        eval_strategy="steps",
        eval_steps=0.1,
        save_strategy="steps",
        save_steps=0.05,
        logging_steps=0.01,
        run_name="mLateOn-medical",
    )

    # 6️⃣ 訓練器與訓練
    trainer = MultiVectorEncoderTrainer(
        model=model,
        args=args,
        train_dataset=train_dataset,
        loss=loss,
        evaluator=dev_evaluator,
    )
    trainer.train()

    # 7️⃣ 儲存並選擇性推送至 Hub
    model.save_pretrained("models/mLateOn-medical/final")
    try:
        model.push_to_hub("mLateOn-medical")
    except Exception:
        logging.error("無法上傳模型:\n" + traceback.format_exc())

if __name__ == "__main__":
    main()
  • 執行時間:單一 RTX 3090 上需 14.5 小時(峰值 17.5 GB VRAM)。
  • 資料效率:100 k 對(約 75 分鐘)即可達成與完整百萬對訓練相差不到 0.012 NDCG@10 的表現。

索引大小與優化

多向量索引較大,因每個詞元產生一個向量(平均 941 個詞元段落產生約 878 個向量)。200 k 段落的原始 fp16 儲存空間約為 45 GB。

詞元聚合

HierarchicalTokenPooling(pool_factor=4) 可將向量數減少 4 倍,損失小於 0.0033 NDCG@10。

from sentence_transformers.multi_vector_encoder.modules import HierarchicalTokenPooling
pooling = HierarchicalTokenPooling(pool_factor=4)
embeddings = model.encode_document(passages, token_pooling=pooling)
  • 向量數減少至 1/4 → 約 11 GB,NDCG@10 ≈ 0.8991。

量化與剪枝(PLAID)

使用 1 位元殘差量化與適度剪枝:

設定 保留向量 索引大小 NDCG@10
1‑bit PLAID,全部向量 100 % 3.37 GB 0.8984
1‑bit PLAID + 剪枝 65 % 2.23 GB 0.8830
1‑bit PLAID + 剪枝 42 % 1.45 GB 0.8642
  • 量化可帶來超過 13 倍的空間壓縮,損失小於 0.02 NDCG,使多向量檢索的儲存空間與密集模型相當。

評估結果

微調後的 multi-vector-encoder/mLateOn-medical 與 50 多種基線模型在 200 k 段落醫療基準上比較(1 k 保留查詢):

模型 類別 NDCG@10 acc@1
multi-vector-encoder/mLateOn-medical 多向量(微調) 0.9139 0.849
lightonai/mLateOn 多向量(零樣本) 0.8520 0.758
lightonai/GTE-ModernColBERT-v1(提升限制) 多向量(零樣本) 0.8502 0.763
Qwen/Qwen3-Embedding-4B 密集(零樣本) 0.7817 0.669
voyageai/voyage-4-nano 密集(零樣本) 0.7563 0.638
BM25 詞法 0.7501 0.641
naver/splade-v3 稀疏(零樣本) 0.6853 0.574
結論: 當文件長度較長時,晚期互動模型表現領先;微調比最強的零樣本模型再提升 +0.062 NDCG@10

實務建議

  1. 從非監督檢查點(例如 lightonai/mLateOn-unsupervised)開始進行領域微調。
  2. 解除文件長度限制 以匹配你的資料;否則在長段落上可能損失最高 0.24 NDCG@10。
  3. 使用 GradCacheCachedMultiVectorMultipleNegativesRankingLoss)在單一 GPU 上實現大有效批次大小。
  4. 套用標點符號跳過清單 可縮小索引約 10 %,且不損失品質。
  5. 投資於索引壓縮(詞元聚合、PLAID 量化)以使多向量儲存空間與密集模型相近。
  6. 即使少量資料(100 k 對)也能產生強大表現,使此方法對許多組織都可行。

額外資源

  • 訓練範例:MIRIAD 醫療微調、MS MARCO 對比與蒸餾、多模態 ColPali、PEFT LoRA 适配器。
  • 文件:安裝、快速入門、使用方式、自訂模型建立、預訓練模型列表、訓練與損失概覽、API 參考。
  • 附帶文章(使用說明):《使用 Sentence Transformers 的多向量(晚期互動)嵌入模型》——涵蓋編碼、索引與部署。

"認為多向量索引過於龐大的反對意見,在正確配置的索引下不成立。" – 感謝 Omar Khattab 提供量化測量。

Sources

相關