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.Dataset 或 DatasetDict,包含查询-文档对(或损失所需的其他格式) |
| 损失 | 批内负样本对比损失(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 或本地)。所需格式取决于损失函数:
- 标签列(
label或score),如果损失需要监督信号。 - 输入列 按顺序排列;第一列为查询,后续列为文档(除非通过
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。 |
实用建议
- 从无监督检查点(如
lightonai/mLateOn-unsupervised)开始进行领域微调。 - 解除文档长度限制 以匹配你的数据;否则在长段落上可能损失高达 0.24 NDCG@10。
- 使用 GradCache(
CachedMultiVectorMultipleNegativesRankingLoss)在单张 GPU 上实现大有效批量。 - 应用标点符号跳过列表 可在不损失质量的情况下将索引缩小约 10%。
- 投入索引压缩(词元池化、PLAID 量化)以使多向量存储接近密集模型水平。
- 即使少量数据(10 万个样本对)也能获得强劲性能,使该方法对众多组织可行。
额外资源
- 训练示例:MIRIAD 医学微调、MS MARCO 对比与蒸馏、多模态 ColPali、PEFT LoRA 适配器。
- 文档:安装、快速入门、使用、自定义模型创建、预训练模型列表、训练与损失概览、API 参考。
- 配套文章(使用指南):《使用 Sentence Transformers 的多向量(延迟交互)嵌入模型》——涵盖编码、索引与服务。
"认为多向量索引过大这一观点,在配置得当的索引面前不成立。" – 感谢 Omar Khattab 提供量化测量数据。
Sources
相关
- Dispatch
- Dispatch
- Dispatch
- Dispatch
- Dispatch