Sentence Transformers 6.0 MultiVectorEncoder:训练与微调指南

TL;DR

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,000 对样本后也能接近检查点性能。

应选择哪个起点?

在 25,000 对 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。
  • 缩放参数:对 MaxSim 得分保留默认值 scale=1.0;使用密集模型默认值(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⁻⁶ → 2e⁻⁴ 的学习率扫描,更高的学习率(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,
)
  • 包含硬干扰段落以避免饱和;博客中将约 19 万个随机训练段落添加到 1 万个黄金集。

完整训练脚本

以下脚本可复现 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️⃣ 加载 100 万个医学 QA 对
    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)。
  • 数据效率:10 万个样本对(约 75 分钟)即可达到与完整百万对训练相差不到 0.012 NDCG@10 的性能。

索引大小与优化

多向量索引更大,因为每个词元都会生成一个向量(平均 941 词元段落约 878 个向量)。20 万个段落的原始 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 位 PLAID,全部向量 100% 3.37 GB 0.8984
1 位 PLAID + 剪枝 65% 2.23 GB 0.8830
1 位 PLAID + 剪枝 42% 1.45 GB 0.8642
  • 量化带来超过 13 倍的压缩,损失小于 0.02 NDCG,使多向量检索的存储成本与密集模型相当。

评估结果

微调后的 multi-vector-encoder/mLateOn-medical 与 50 多个基线模型在 20 万个段落的医学基准测试(1,000 个保留查询)上的对比:

模型 类型 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. 即使少量数据(10 万个样本对)也能获得强劲性能,使该方法对众多组织可行。

额外资源

  • 训练示例:MIRIAD 医学微调、MS MARCO 对比与蒸馏、多模态 ColPali、PEFT LoRA 适配器。
  • 文档:安装、快速入门、使用、自定义模型创建、预训练模型列表、训练与损失概览、API 参考。
  • 配套文章(使用指南):《使用 Sentence Transformers 的多向量(延迟交互)嵌入模型》——涵盖编码、索引与服务。

"认为多向量索引过大这一观点,在配置得当的索引面前不成立。" – 感谢 Omar Khattab 提供量化测量数据。

Sources

相关