用 Sentence Transformers 训练与微调多向量嵌入模型
Training and Finetuning Multi-Vector Embedding Models with Sentence Transformers
Sentence Transformers v6.0 新增第四种模型类型 MultiVectorEncoder,支持 ColBERT 风格的后交互检索,并配套完整训练流程。
文章量化了不同起点的领域适应差距,未监督预训练 checkpoint 反超已微调版本,并把长文档截断列为比架构更影响 NDCG 的因素。
Sentence Transformers 是一个 Python 库,用于使用和训练嵌入向量与重排序模型,应用范围广泛,例如检索增强生成、语义搜索、语义文本相似度等。其 v6.0 更新引入了第四种模型类型:MultiVectorEncoder,用于 ColBERT 风格的后期交互检索,并配套了完整的训练方法。在这篇博文中,我将向你展示如何使用它来微调一个多向量模型,使其在你的数据上超越通用检索器。该方法还可以从零开始训练出强大的全新多向量模型。以下所有内容都运行在 pip install -U "sentence-transformers[train]" 上。微调多向量模型涉及多个组件:模型本身、数据集、损失函数、训练参数、评估器和训练器类。我将逐一介绍这些组件,并配以实际示例,说明如何用它们来微调强大的多向量模型。
最后,在评估部分,我将向你展示,我微调后的 multi-vector-encoder/mLateOn-medical 模型——伴随这篇博文在单张 RTX 3090 上训练了 14.5 小时——在我的医学检索评估中轻松超越了所有我能找到的通用检索模型:无论是稠密、稀疏、词法还是多向量模型,一概如此。
如果你有兴趣微调稠密嵌入向量模型、稀疏嵌入向量模型或重排序模型,那么不妨阅读我先前的博文:训练与微调嵌入向量模型、训练与微调稀疏嵌入向量模型 以及 训练与微调重排序模型。
这篇博客文章讲的是训练多向量模型。如果你想了解如何使用它们,从加载、编码到在向量数据库中建立索引,请参阅配套的使用 Sentence Transformers 的多向量(后期交互)嵌入向量模型博客文章。
目录
什么是多向量模型?
稠密嵌入向量模型会把整段文本压缩成单个向量,相似度就是两个这样的摘要之间的一次点积。多向量模型(也称为后期交互模型或 ColBERT 风格模型)则跳过了这种压缩。它为每个 token 保留一个小向量,并用 MaxSim 算子对查询与文档打分,即每个查询 token 找到与其最匹配的文档 token,然后将得分求和。token 级别的匹配恰好保留了单个向量不得不平均掉的细粒度信号,这通常意味着更强的检索能力,代价是更大的索引。
配套的多向量嵌入向量模型博客文章详细介绍了架构、编码、打分和索引,所以这部分我会讲得简短一些,直接进入训练环节。
为什么要微调?
微调多向量模型能显著提升其在你的特定领域上的检索性能:网络搜索、法律取证、代码搜索和科学文献综述之间的词汇、查询风格以及相关性的定义都各不相同。由于查询和文档是逐 token 匹配的,多向量模型能够捕捉到单向量模型往往会平均掉的细粒度领域信号,而且即使只有适量的领域内微调数据,它们也能有非常好的表现。
除此之外,大多数已发布的检索模型都是为短段落配置的。经典的 ColBERT 检查点会将文档截断到 180 或 300 个 token,而许多流行的稠密模型则截断到 256 或 512 个 token,因为它们基于 MS MARCO 风格的训练数据很少超过这个长度。如果你的文档很长,这些模型会在打分之前悄悄丢弃每篇文档的大部分内容。在我那次段落平均长度为 941 个 token 的医学评测中,我测得这种截断带来的损失高达 0.24 NDCG@10,远超任何模型架构之间的差异。当你训练自己的模型时,你可以根据你的数据需要来配置文档长度。
LightOn 在代码检索上也遇到了同样的情形,通用的 LateOn 不够用,于是他们训练了 LateOn-Code。你的领域,无论是医学、法律、金融,还是你公司的内部文档,都不会有官方模型。这篇博文将向你展示如何自己动手构建,只需几个小时,在一块消费级 GPU 上即可完成。
训练组件
训练 MultiVectorEncoder 模型涉及以下组件:
- 模型:要微调的模型,或从零构建的架构。
- 数据集:用于训练和评估的数据。
- 损失函数:衡量模型性能并指导优化过程的函数。
- 训练参数(可选):影响训练性能、跟踪和调试的参数。
- 评估器(可选):用于在训练前、训练中或训练后评估模型的类。
- Trainer:将所有训练组件整合在一起。
让我们逐一深入了解每个组件。
模型
多向量训练让你真正可以选择起点,而它的重要性可能超出你的预期。
微调现有的多向量模型
如果你想进一步微调现有的多向量模型,完全不必担心架构问题:
from sentence_transformers import MultiVectorEncoder
# Loading in fp32 is preferred for training if your memory can handle it
model = MultiVectorEncoder(
"lightonai/mLateOn-unsupervised",
model_kwargs={"torch_dtype": "float32"},
processor_kwargs={"model_max_length": 8192}, # the tokenizer-level token limit
)
该 checkpoint 自带一整套配置:它的查询与文档标记 token、它的投影头、它的评分 skiplist。进行微调时,通常你会希望保留所有这些,只改动你的数据所要求的部分。首先要检查的是长度配置,因为许多已发布的 checkpoint 将文档上限设为 180 到 512 个 token(参见 为什么要微调?),而我的医学段落长达 1,400 个 token。mLateOn 系列已经支持骨干模型完整的 8192 token 上下文,但如果你起始的 checkpoint 带有上限,就把它们解除:
# Let the model read full documents instead of the caps it was trained with,
# e.g. GTE-ModernColBERT-v1 ships with query_length=48 and document_length=300
model[0].query_length = None
model[0].document_length = None
在未设置各任务上限的情况下,截断会回退到 tokenizer 的 model_max_length,这就是为什么我在上面的加载时配置了该限制。
我还做了一处改动,添加了一个标点 skiplist,将标点 token 从文档侧评分和存储中排除。在一次 4 路消融实验(无、标点、停用词、两者兼有)中,它在质量上小幅胜出,并且在此数据上免费将文档索引缩小了 9.6%:
import string
# model[2] is the MultiVectorMask module
model[2].skiplist_words = list(string.punctuation)
model[2].resolve_with_tokenizer(model.tokenizer) # token ids are cached, so re-resolve after changing
从基础 transformer 构建一个
你也可以将 MultiVectorEncoder 指向任意基础 transformer,系统会为你附加一个全新的、随机初始化的 token 级投影:
from sentence_transformers import MultiVectorEncoder
model = MultiVectorEncoder("answerdotai/ModernBERT-base", model_kwargs={"torch_dtype": "float32"})
# MultiVectorEncoder(
# (0): Transformer({..., 'architecture': 'ModernBertModel'})
# (1): Dense({'in_features': 768, 'out_features': 128, 'bias': False, ...})
# (2): MultiVectorMask({'skiplist_words': [], 'skiplist_tasks': ['document'], ...})
# (3): Normalize({...})
# )
这就是经典的 ColBERT 流水线:一个 Transformer 生成上下文化的 token 嵌入向量,一个 token 级的 Dense 将每个嵌入向量投影到 128 维,一个 MultiVectorMask 决定打分时哪些 token 计入,以及一个 token 级的 Normalize。投影初始是随机的,因此该模型在可用之前需要经过训练。有趣的是,这在强大的稠密嵌入向量骨干网络上同样有效。在我的实验中,在 Alibaba-NLP/gte-modernbert-base 上全新初始化的投影,仅凭投影层和 25k 训练对,就从零开始达到了与现有检查点起点相差 0.03 以内的水平。
经典的 ColBERT token 化技巧([MASK] 查询扩展、[Q] / [D] 前缀 token、文档长度上限、标点跳过列表)默认全部关闭且可配置。完整列表参见 Creating Custom Models。顺便一提,我在自己的领域微调中测试了 [MASK] 查询扩展的四种配置,没有一种带来可测量的差异,所以不必强求套用经典配方。
你应该选择哪个起点?
我在准备这篇博文时直接做了测量,取六个起点,用完全相同的配方在来自 MIRIAD 的 25k 医学问答对上分别训练,然后在 1,000 个留出问题上针对 50,000 篇文档的语料库进行评估:
| 起点 | 零样本 NDCG@10 | 25000 对之后 | 差值 |
|---|---|---|---|
| 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 | - |
这个结果让我感到意外,而且它在两个模型家族中都得到了复现。*这些 -unsupervised checkpoint 对新领域的适应能力远超它们已完成的同类版本,尽管起点更低,却实现了反超。这些 checkpoint 处于大规模对比预训练之后、但在通用检索的监督微调之前,因此它们携带了全部后期交互结构,却没有那些通用目的的调优——而领域训练随后不得不把这些调优抹掉。相比之下,已完成的 checkpoint 在我尝试的每一个学习率下都几乎纹丝不动,甚至出现退化。
所以,如果你喜欢的模型家族发布了监督前的 checkpoint,就从那里开始。如果没有,那么在强大的检索预训练骨干网络上重新初始化投影头是紧随其后的次优选择。从完全训练完成的 checkpoint 继续,是领域适应中最弱的选择,尽管它感觉上最自然。
数据集
MultiVectorEncoderTrainer 使用 datasets.Dataset 或 datasets.DatasetDict 实例进行训练和评估。你可以从 Hugging Face Datasets Hub 加载数据,也可以使用你喜欢的任意格式的本地数据(例如 CSV、JSON、Parquet、Arrow 或 SQL)。
注意: 许多可直接配合 Sentence Transformers 使用的公开数据集已在 Hugging Face Hub 上被标记为 sentence-transformers,因此你可以在 https://huggingface.co/datasets?other=sentence-transformers 上轻松找到它们。不妨浏览一下这些数据集,找到可能对你的任务、领域或语言有用的现成数据集。
Hugging Face Hub 上的数据
你可以使用 load_dataset 函数从 Hub 上的数据集加载数据:
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: 4467542
})
"""
这是我在本篇博文中将要训练所用的数据集:来自 MIRIAD 的 440 万个医学问题,每个问题都配有一段包含其答案的来源段落(平均 941 个 token)。像这样简单的(查询,相关段落)对,是为你自己领域收集检索训练数据时最容易获取的,而且正如你将看到的,它们就是你所需要的一切。
本地数据
你也可以使用 load_dataset 来加载常见文件格式的本地数据:
from datasets import load_dataset
dataset = load_dataset("csv", data_files="my_file.csv")
# or
dataset = load_dataset("json", data_files="my_file.json")
如果你的本地数据需要预处理,你可以使用 datasets.Dataset.from_dict 以字典列表的形式初始化数据集:
from datasets import Dataset
queries = []
documents = []
# Open a file, perform preprocessing, filtering, cleaning, etc.
# and append to the lists
dataset = Dataset.from_dict({
"query": queries,
"document": documents,
})
数据集格式
重要的是,你的数据集格式要与损失函数相匹配(或者你选择一个与数据集格式相匹配的损失函数)。验证数据集格式是否与损失函数兼容,涉及两个步骤:
- 如果你的损失函数需要标签根据损失函数概览表,那么你的数据集必须有一个名为 "label" 或 "score" 的列。该列会自动被作为标签。
- 所有未命名为“label”或“score”的列都被视为输入根据损失概览表格。剩余列的数量必须与你所选损失函数的有效输入数量相匹配。这些列的名称无关紧要,只有顺序重要.
在此基础上还有两个多向量特有的约定:
- 位置查询与文档分配:无论列名是什么,第一列都会被嵌入为查询,其后所有列都会被嵌入为文档。这一默认行为可以通过标准的
router_mapping训练参数按列进行覆盖。 - 知识蒸馏格式:每个候选文档一列,即
(query, document_1, ..., document_N, scores),其中scores是每行 N 个教师分数的列表。对于将查询和文档 ID 与单独的文本数据集(例如 lightonai/ms-marco-en-bge)一起存储的 KD 数据集,你可以使用resolve_ids来即时将 ID 解析为文本。
损失函数
损失函数量化模型在给定一批数据上的表现好坏,使优化器能够更新模型权重,从而产生更有利(即更低)的损失值。适合你任务的损失函数取决于你拥有的数据以及你想要达成的目标。你可以在 损失函数概览 中找到完整的选项列表。
对于问答或问题-段落对这类常见场景,主力方法是使用 MultiVectorMultipleNegativesRankingLoss 进行批内负样本训练,其中批次中的每个其他文档都作为每个查询的负样本。更大的批次意味着更多的负样本和更强的训练效果,因此实践中你会需要它的 GradCache 变体 CachedMultiVectorMultipleNegativesRankingLoss,它将有效批次大小与 GPU 上能容纳的大小解耦:
from sentence_transformers import MultiVectorEncoder
from sentence_transformers.multi_vector_encoder.losses import CachedMultiVectorMultipleNegativesRankingLoss
model = MultiVectorEncoder("lightonai/mLateOn-unsupervised", model_kwargs={"torch_dtype": "float32"})
loss = CachedMultiVectorMultipleNegativesRankingLoss(
model=model,
mini_batch_size=16, # how many documents to encode per chunk: bounds memory, not quality
)
mini_batch_size 参数通过按此大小分块编码文档来限制内存,而有效对比批大小(在我下面的运行中是 128,在我的消融实验中更大的批次没有带来任何进一步收益)仍可自由选择。GradCache 保证无论分块大小如何结果都完全相同,因此对于较小的 GPU 可以降低该值,代价仅是墙钟时间。当你的文档长度变化很大时,可以考虑它的同类 mini_batch_num_tokens,它按总 token 预算而非文档数量来打包每个分块,因此一个包含异常长文档的分块永远不会让你的内存飙升(我的 mini_batch_size=16 在每篇文档约 940 个 token 时对应 mini_batch_num_tokens=15_000)。
多向量特有的一个陷阱是,对比损失默认使用 scale=1.0,而稠密嵌入的对应默认值是 scale=20.0。那个 20.0 之所以存在,是因为余弦相似度是 [-1, 1] 范围内的单个值,对于尖锐的 softmax 来说范围太窄。而 MaxSim 分数则是为每个查询 token 汇总一个最佳匹配相似度,因此它的范围大致是 [0, query_length]:一个 32 token 的查询最高可以得 32 分。所以不要把 scale=20.0 从稠密训练脚本里照搬过来,否则它会让 softmax 饱和并杀死你的梯度。
关于从更强的教师模型进行知识蒸馏——这也是最强的通用 late-interaction 模型的训练方式——请参见 MultiVectorDistillKLDivLoss 以及 Training Overview 文档中的 Knowledge Distillation 标签页。
训练参数
你可以使用 MultiVectorEncoderTrainingArguments 类来自定义训练过程。该类让你可以调整会影响训练速度的参数,并帮助你理解训练过程中发生了什么。
关于最实用的训练参数的更多信息,请查看 Multi-Vector Encoder > Training Overview > Training Arguments。值得一读,以便充分利用你的训练。
下面是一个示例,使用的是我实际训练运行中的取值:
from sentence_transformers import MultiVectorEncoderTrainingArguments
from sentence_transformers.base.sampler import BatchSamplers
args = MultiVectorEncoderTrainingArguments(
# Required parameter:
output_dir="models/mLateOn-medical",
# Optional training parameters:
num_train_epochs=1,
per_device_train_batch_size=128, # the effective contrastive batch, thanks to GradCache
per_device_eval_batch_size=16,
learning_rate=1e-4,
warmup_steps=0.05,
prompts={"question": "[Q] ", "passage_text": "[D] "}, # the checkpoint's markers, keyed by training column
fp16=False, # Set to True if you have a GPU that supports FP16
bf16=True, # Set to True if you have a GPU that supports BF16
batch_sampler=BatchSamplers.NO_DUPLICATES, # in-batch negatives benefit from no duplicates
# Optional tracking/debugging parameters:
eval_strategy="steps",
eval_steps=0.1,
save_strategy="steps",
save_steps=0.05,
logging_steps=0.01,
run_name="mLateOn-medical", # Will be used in e.g. Trackio, W&B, etc.
)
其中几项值得说明一下:
prompts:训练不会自动应用存储在模型中的提示词,因此需要将它们显式映射到你的训练列上。在这里,也就是该 checkpoint 的[Q]标记对应问题列,[D]对应段落列,从而保持训练与推理一致。max_length(故意不设置):该参数仅在 训练时限制 tokenization,适用于你希望训练成本低于模型完整服务长度的情况。我测量了这种捷径在这份数据上的代价。以 512 tokens 训练损失了约 0.015 NDCG@10,换来了约 2 倍的速度,而且这一差距并不会随着数据增多而缩小,因为模型根本看不到被截断的内容。除非你更需要速度提升而非质量,否则不要设置它,让训练与推理保持一致。learning_rate=1e-4:在从 5e-6 到 2e-4 进行了一轮扫描后,我发现这个高于通常水平的学习率效果最好。
评估器
为了在训练过程中跟踪模型的表现,你可以向训练器传入一个 eval_dataset 来评估损失,但具体的检索指标信息量要大得多。Sentence Transformers 为多向量模型内置了以下评估器:
| 评估器 | 所需数据 |
|---|---|
MultiVectorInformationRetrievalEvaluator | 查询、语料库以及相关文档映射 |
MultiVectorNanoBEIREvaluator | 无需数据 |
MultiVectorTripletEvaluator | (锚点, 正例, 负例) 三元组 |
MultiVectorRerankingEvaluator | {'query': '...', 'positive': [...], 'negative': [...]} 字典列表 |
MultiVectorDistillationEvaluator | 带有候选文档和教师分数的查询 |
对于领域微调而言,真正重要的是用你自己留出的数据构建的 MultiVectorInformationRetrievalEvaluator。关于构建它有一个建议:语料库要足够难,这样才能把不同模型区分开来。在我的案例中,MIRIAD 的问题是从它们各自的源段落生成的,这使得检索异常简单。仅针对那 1 万条金标段落,几乎每个模型的 NDCG@10 都超过了 0.97。如果你的评测也这样饱和了,就加入 干扰段落(我使用的是训练集中去重后的段落),直到分数拉开差距:
from datasets import load_dataset
from sentence_transformers.multi_vector_encoder.evaluation import MultiVectorInformationRetrievalEvaluator
dataset = load_dataset("tomaarsen/miriad-4.4M-split")
# Gold: 1,000 evaluation questions, each mapping to its own passage, with the
# eval split's full ~10k unique passages as the initial corpus
corpus = {}
queries = {}
relevant_docs = {}
passage_to_id = {}
for idx, row in enumerate(dataset["eval"]):
if row["passage_text"] not in passage_to_id:
passage_to_id[row["passage_text"]] = f"p{len(passage_to_id)}"
corpus[passage_to_id[row["passage_text"]]] = row["passage_text"]
if idx < 1_000:
queries[f"q{idx}"] = row["question"]
relevant_docs[f"q{idx}"] = {passage_to_id[row["passage_text"]]}
# Distractors: unique train passages that make the haystack realistic
seen = set(passage_to_id)
for row in dataset["train"]:
if len(corpus) >= 200_000:
break
if row["passage_text"] not in seen:
seen.add(row["passage_text"])
corpus[f"d{len(corpus)}"] = row["passage_text"]
evaluator = MultiVectorInformationRetrievalEvaluator(
queries=queries,
corpus=corpus,
relevant_docs=relevant_docs,
name="miriad-dev",
batch_size=16,
)
# results = evaluator(model)
训练器
MultiVectorEncoderTrainer 是所有前面组件汇聚到一起的地方。以下是训练 multi-vector-encoder/mLateOn-medical(即引言中提到的那个模型)的完整脚本:
import logging
import string
import 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.evaluation import MultiVectorInformationRetrievalEvaluator
from sentence_transformers.multi_vector_encoder.losses import CachedMultiVectorMultipleNegativesRankingLoss
logging.basicConfig(format="%(asctime)s - %(message)s", datefmt="%Y-%m-%d %H:%M:%S", level=logging.INFO)
def main():
# 1. Load the starting checkpoint: contrastively pretrained, not yet supervised
# Loading in fp32 is preferred for training if your memory can handle it
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 finetuned on MIRIAD medical retrieval",
),
)
# 2. Lift the per-task length caps so training and inference see full medical passages
model[0].query_length = None
model[0].document_length = None
# 3. Skip punctuation tokens during scoring: a small quality win and a 9.6% smaller index
model[2].skiplist_words = list(string.punctuation)
model[2].resolve_with_tokenizer(model.tokenizer)
# 4. Load 1 million medical question-passage pairs
train_dataset = load_dataset("tomaarsen/miriad-4.4M-split", split="train").select(range(1_000_000))
# 5. In-batch negatives with GradCache: large effective batch, memory-bounded chunks
loss = CachedMultiVectorMultipleNegativesRankingLoss(model=model, mini_batch_size=16)
# 6. A light dev evaluator to watch progress during training: 500 held-out questions
# against the eval split's ~10k unique passages. The full 200k protocol runs afterwards.
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:
passage_to_id[row["passage_text"]] = f"p{len(passage_to_id)}"
corpus[passage_to_id[row["passage_text"]]] = row["passage_text"]
if idx < 500:
queries[f"q{idx}"] = row["question"]
relevant_docs[f"q{idx}"] = {passage_to_id[row["passage_text"]]}
dev_evaluator = MultiVectorInformationRetrievalEvaluator(
queries=queries, corpus=corpus, relevant_docs=relevant_docs, name="miriad-dev", batch_size=16
)
# 7. Training arguments, as discussed above
run_name = "mLateOn-medical"
args = MultiVectorEncoderTrainingArguments(
output_dir=f"models/{run_name}",
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, # Set to True if you have a GPU that supports FP16
bf16=True, # Set to True if you have a GPU that supports BF16
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=run_name,
)
# 8. Create a trainer & train
trainer = MultiVectorEncoderTrainer(
model=model,
args=args,
train_dataset=train_dataset,
loss=loss,
evaluator=dev_evaluator,
)
trainer.train()
# 9. Save the trained model
model.save_pretrained(f"models/{run_name}/final")
# 10. (Optional) Push it to the Hugging Face Hub
try:
model.push_to_hub(run_name)
except Exception:
logging.error(f"Error uploading model to the Hugging Face Hub:\n{traceback.format_exc()}")
if __name__ == "__main__":
main()
这就是完整的配方:一个预监督检查点、一百万对领域数据、批内负样本、完整文档长度,以及比通常更高的学习率。这次运行在我的单块 RTX 3090 上耗时 14.5 小时,峰值显存占用 17.5 GB,而其中每一个选择都是经过实测比较后的胜出者,而非猜测。
对于预算更有限的读者,我的扩展实验表明,10 万对数据(训练 75 分钟)与完整的百万对训练相比,NDCG@10 仅差 0.012。大部分收益都来自第一个小时。
回调
MultiVectorEncoder 训练器支持多种 transformers.TrainerCallback 子类,包括:
WandbCallback,用于在安装了wandb时将训练指标记录到 W&BTensorBoardCallback,用于在可访问tensorboard时将训练指标记录到 TensorBoardCodeCarbonCallback,用于在安装了codecarbon时跟踪训练期间的碳排放
通过 report_to 训练参数启用这些功能,例如 report_to=["wandb", "codecarbon"],并安装所需的依赖项。它默认为 "none",而 report_to="all" 会激活所有已安装依赖项的集成。
有关这些回调以及如何创建自己的回调的更多信息,请参阅 Transformers Callbacks 文档。
多数据集训练
通常,表现最佳的通用模型会同时在多个数据集上进行训练。然而,由于每个数据集的格式各不相同,这种方法可能颇具挑战性。幸运的是,MultiVectorEncoderTrainer 允许你在多个数据集上进行训练,而无需统一格式。此外,它还提供了为每个数据集应用不同损失函数的灵活性。以下是同时使用多个数据集进行训练的步骤:
- 使用一个由
datasets.Dataset实例组成的字典(或一个datasets.DatasetDict)作为train_dataset(也可选地作为eval_dataset)。 - (可选)使用一个损失函数字典,将数据集名称映射到损失函数。仅当你希望为不同数据集使用不同损失函数时才需要。
每个训练/评估批次只会包含来自其中一个数据集的样本。从多个数据集中采样批次的顺序由MultiDatasetBatchSamplers枚举定义,该枚举可以传递给MultiVectorEncoderTrainingArguments通过multi_dataset_batch_sampler。有效选项为:
MultiDatasetBatchSamplers.ROUND_ROBIN:从每个数据集中轮流采样,直到其中一个数据集被耗尽。采用这种策略,很可能并非每个数据集中的所有样本都会被使用,但每个数据集被采样的机会是均等的。MultiDatasetBatchSamplers.PROPORTIONAL(默认):按照每个数据集的大小比例从中采样。采用这种策略,每个数据集中的所有样本都会被使用,且较大的数据集被采样的频率更高。
评估
为了弄清微调后的模型处于什么水平,我在 MIRIAD 评测集上,针对四个架构家族的 50 多种检索模型配置对它进行了评估,评测集的构建方式与上文 Evaluator 一节完全一致,使用 1,000 道留出的医学问题检索 200,000 个唯一段落(10k 个金标准段落隐藏在来自训练划分的 190k 个去重干扰项之中)。该语料库的规模是 Which starting point should you pick? 中 50,000 段落语料库的四倍,因此两张表之间的分数不可直接比较。
核心结果如下,完整表格见下方可折叠部分:
| 模型 | 家族 | NDCG@10 |
|---|---|---|
| multi-vector-encoder/mLateOn-medical(我的) | 多向量,微调 | 0.9139 |
| lightonai/mLateOn | 多向量,零样本 | 0.8520 |
| lightonai/GTE-ModernColBERT-v1(上限已解除) | 多向量,零样本 | 0.8502 |
| Qwen/Qwen3-Embedding-4B | 稠密,零样本 | 0.7817 |
| voyageai/voyage-4-nano | 稠密,零样本 | 0.7563 |
| BM25 | 词汇 | 0.7501 |
| naver/splade-v3 | 稀疏,零样本 | 0.6853 |
微调后的模型位居榜首,以 +0.062 NDCG@10 的成绩击败了任何架构下最强的零样本模型。换句话说,最强的零样本模型在 75.8% 的查询中将正确段落作为第一个命中结果返回,而微调后的模型则在 84.9% 的查询中做到这一点,将排名第一的错误率削减了超过三分之一。
架构模式同样清晰,榜单前列清一色是后期交互。在长文档上,每个 token 一个向量胜过每个文档一个向量,即使在训练数据和骨干网络都匹配的情况下也是如此。DenseOn 和 LateOn 共享训练数据和架构,仅头部不同,而后期交互的兄弟模型以 +0.12 胜出,多语言配对(mDenseOn 和 mLateOn)以 +0.13 复现了这一结果。规模也无法拯救单向量方案。Qwen3-Embedding-4B是其中最强的稠密模型,活跃(非嵌入)参数量约为我的 33 倍,但仍落后 0.13,而 8B 版本的得分还低于 4B。
BM25 的表现也出人意料地好,击败了所有稀疏模型、所有截断上限的多向量模型,以及除三个稠密模型之外的所有模型:数十亿参数的 Qwen3-Embedding-4B和 8B,以及 voyage-4-nano,后者读取其完整的 32k token 上下文,仅以 0.006 的微弱优势胜出。不过别指望这能迁移到你自己的数据上。MIRIAD 的问题是从段落生成的,因此查询与其黄金段落之间的词汇重叠远大于典型检索场景,而 BM25 不受限的上下文长度让它能利用所有这些重叠词,而大多数神经网络检查点则会截断。BM25 基线成本低廉且始终值得运行,只是别指望这个差距。
完整榜单一览,按得分排序,并按架构家族着色。
点击查看完整评测表
标记为 @N 的模型在评估时将其文档长度上限提升至 N tokens,因为它们原生的上限(180 至 512 tokens)否则会截断平均长度为 941 tokens 的段落。对于每一个多向量模型,这一提升相比按原样提供的结果在 NDCG@10 上带来了 +0.08 至 +0.24 的增益,甚至稠密的 DenseOn 也因同样的处理获得了 +0.03 的提升。
请注意,这并不意味着 multi-vector-encoder/mLateOn-medical 在所有领域都是最强的模型。它只是在我的领域中最强。这完全没问题,因为我只需要这个模型在我的数据上表现良好。
不要低估在你的领域上微调多向量模型的力量。在一块消费级 GPU 上花费十四个半小时,就产出了一个通用检索器在这份数据上无法企及的模型,而且配方只是一个脚本,没有教师模型,也没有挖掘的负样本!
优化索引
对多向量检索的合理反对意见是索引大小,而这个领域几乎是最糟糕的情况。每个 token 存储一个向量,我的模型每段需要约 878 个向量,因此 20 万段的语料库在 fp16 下大约占用 45 GB,而稠密模型所需的远低于 1 GB。文档长度正是造成这一差距如此之大的原因。配套文章中的 Natural Questions 段落平均每段约 125 个 token 向量,少了七倍,因此短段落语料库的起始索引远小于这个语料库。HierarchicalTokenPooling 模块正是通过聚类每个文档的 token 嵌入向量并存储聚类均值来压缩这一点,大致保留 1 / pool_factor 的向量:
from sentence_transformers.multi_vector_encoder.modules import HierarchicalTokenPooling
pooling = HierarchicalTokenPooling(pool_factor=4)
document_embeddings = model.encode_document(passages, token_pooling=pooling)
我是在完成后的模型上事后测量的,没有进行池化感知训练,而在长文档上它的成本低得惊人。
这些实心点表示未压缩的嵌入向量,因此每个系列都以相同方式计数,并用精确搜索进行评分。不过,你不会以这种方式部署其中任何一个。稠密索引通常使用 int8 或二值量化并配合重打分,稀疏索引会压缩其倒排列表,而多向量索引则使用 PLAID 风格的残差压缩。不要把那些点理解为你需要购买的磁盘容量,而应理解为相对存储成本。
Token 池化是那条实线。将向量数量减半会损失 0.0033 NDCG@10,且 rank-1 准确率不受影响;而只保留四分之一,在 11.2 GB 下仍能得分 0.8991。这条曲线还在继续(我测到了向量的十分之一,仍为 0.8765),但一旦量化方案摆上台面,就没有太多理由把池化推到那么极端,这正是下面那条虚线所要说明的。
虚线代表真实部署可能的样子。我给了 Omar Khattab 该模型和基准的早期访问权限,他用 fast-plaid 在 1-bit 残差量化下测量了这些配置,使用紧凑的 17-bit 质心 id 和 18-bit 文档 id,而非其通常的未打包 64-bit 整数,外加文档侧剪枝:
| 配置 | 保留向量数 | 索引 | 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 倍,NDCG@10 为 0.0155。这比池化曲线上的任何位置都是更划算的取舍。量化会缩小每个向量,而池化和剪枝则减少你保留的向量数量,因此二者可以叠加,而量化是应当优先采用的那一个。再往前推,最后一行落在 1.45 GB,比Qwen3-Embedding-8B 的 fp16 嵌入向量(1.64 GB)更小,同时得分还高出 0.0895。那种认为多向量索引太大的反对意见,在一个配置得当的索引面前是站不住脚的。
这里的剪枝很朴素,只是为了确立 token 削减在量化之上同样有效,所以请把最后两行看作下限,而非前沿。如果你完全不想手动调优量化,配套文章的 Indexing 一节涵盖了 fast-plaid、Qdrant、Weaviate 和 Vespa。
多向量检索的成本完全取决于其索引。该语料库的原始嵌入向量为 45 GB,而配置得当的索引在准确率几乎相同的情况下至少小 7 倍。索引值得你像对待 checkpoint 一样给予同等关注。
致谢
感谢 Omar Khattab 在 优化索引 中测量了量化与剪枝后的索引配置,并围绕后期交互索引成本展开了讨论。
附加资源
训练示例
这些页面包含带有解释的训练示例以及训练脚本的链接。你可以用它们来熟悉多向量训练循环:
- MIRIAD:面向医疗检索的领域特定训练,是本博文方案更早期、更简单的近亲
- MS MARCO:对比学习与知识蒸馏方案
- 多模态:ColPali 风格的视觉文档检索训练
- PEFT 适配器:使用 LoRA 进行参数高效微调
文档
如需进一步学习,你还可以探索以下关于 Sentence Transformers 的资源:
此外,还有一个进阶页面或许你会感兴趣:
以及配套博文,涵盖关于使用这些模型的方方面面:
来源:Hugging Face:Blog(RSS) · huggingface.co