Signal Desk
Hugging Face:Blog(RSS)AI 中文译文 · 待人工复核

使用 Sentence Transformers 训练和微调多向量嵌入模型

英文标题:Training and Finetuning Multi

本文详细介绍了如何使用 Sentence Transformers 训练和微调多向量嵌入模型,并展示了在医学检索任务中,微调后的模型显著优于通用检索模型。

正文研究价值 88 / 100

英文原文我的研究区

为什么值得关注

多向量模型在长文档检索中表现优异,但通用模型往往因截断而损失性能。通过微调,即使使用消费级 GPU 和少量数据,也能训练出领域专用的高性能模型,大幅提升检索准确率。这为医疗、法律、金融等专业领域的搜索应用提供了可行方案,降低了定制化检索模型的门槛。

核心要点

  1. 多向量模型为每个词元保留向量,通过 MaxSim 打分,能捕捉细粒度信号,但索引更大。
  2. 微调多向量模型能显著提升领域特定性能,尤其适合长文档,避免截断损失。
  3. 推荐使用预监督检查点作为起点,比成品检查点适应新领域的能力更强。
  4. 使用 CachedMultiVectorMultipleNegativesRankingLoss 损失函数,结合批内负样本和 GradCache,可高效训练。
  5. 训练时设置完整文档长度,避免截断,并使用比平时更高的学习率(如 1e-4)。
  6. 索引优化(如量化和池化)可大幅减小索引大小,同时保持高准确率。

微调多向量模型涉及几个组成部分:模型本身、数据集、损失函数、训练参数、评估器和训练器类。我将逐一介绍这些组件,并附上如何使用它们微调出强大多向量模型的实用示例。

最后,在评估部分,我会向你展示我微调的 multi-vector-encoder/mLateOn-medical 模型——它是在写这篇博客的同时,用一块 RTX 3090 显卡训练了 14.5 小时得到的——在我的医学检索评估中轻松超越了所有我能找到的通用检索模型,无论是密集向量、稀疏向量、词法匹配还是多向量模型。

如果你对微调密集嵌入模型、稀疏嵌入模型或重排序模型感兴趣,可以阅读我之前写的 训练和微调嵌入模型训练和微调稀疏嵌入模型训练和微调重排序模型 博客文章。

这篇博客讲的是训练多向量模型。如果你想学习如何使用它们,从加载、编码到在向量数据库中建立索引,请参阅配套的 使用 Sentence Transformers 的多向量(后期交互)嵌入模型 博客文章。

目录

什么是多向量模型?

密集嵌入模型把整段文本压缩成一个向量,相似度就是两个这样的摘要之间的一个点积。多向量模型(也叫后期交互或 ColBERT 风格模型)跳过了这种压缩。它为每个词元(token,即文本中的最小单位,比如一个词或一个子词)保留一个小向量,并用 MaxSim 算子来给查询和文档打分:每个查询词元找到它最匹配的文档词元,然后把所有分数加起来。词元级别的匹配保留了单个向量不得不平均掉的细粒度信号,这通常意味着更强的检索能力,代价是索引更大。

配套的 多向量嵌入模型 博客详细介绍了架构、编码、打分和索引,所以我这里就简短带过,直接进入训练部分。

为什么要微调?

微调多向量模型能显著提升它们在你特定领域的检索性能:网络搜索、法律检索、代码搜索和科学文献综述之间的词汇、查询风格和相关性概念都不同。因为查询和文档是逐词元匹配的,多向量模型能捕捉到单向量模型往往会平均掉的细粒度领域信号,而且它们对即使是少量的领域内微调数据也响应得很好。

此外,大多数已发布的检索模型是为短段落配置的。经典的 ColBERT 检查点把文档截断在 180 或 300 个词元,许多流行的密集模型截断在 256 或 512 个词元,因为它们的 MS MARCO 风格训练数据很少超过这个长度。如果你的文档很长,这些模型在打分之前就会默默丢弃每篇文档的大部分内容。在我的医学评估中,段落平均有 941 个词元,我测出这种截断会损失高达 0.24 的 NDCG@10(一种衡量检索排序质量的指标),这比任何模型架构之间的差异都要大得多。当你训练自己的模型时,你可以配置你的数据需要的文档长度。

LightOn 公司在代码检索中也遇到了同样的问题,通用的 LateOn 不够用,于是他们训练了 LateOn-Code。你的领域——无论是医学、法律、金融还是你公司的内部文档——不会有官方发布的模型。这篇博客文章告诉你如何自己构建它,只需几个小时,用一块消费级 GPU 就能完成。

训练组件

训练 MultiVectorEncoder 模型涉及以下组件:

  • 模型:要微调的模型或要全新构建的架构。
  • 数据集:用于训练和评估的数据。
  • 损失函数:衡量模型性能并指导优化过程的函数。
  • 训练参数(可选):影响训练性能、跟踪和调试的参数。
  • 评估器(可选):用于在训练前、训练中或训练后评估模型的类。
  • 训练器:把所有训练组件整合在一起。

让我们仔细看看每个组件。

模型

多向量训练给了你真正的起点选择,而这个选择比你想象的更重要。

微调现有的多向量模型

如果你想进一步微调一个现有的多向量模型,你完全不需要担心架构问题:

from sentence_transformers import MultiVectorEncoder

# 如果内存允许,训练时优先以 fp32 格式加载
model = MultiVectorEncoder(
    "lightonai/mLateOn-unsupervised",
    model_kwargs={"torch_dtype": "float32"},
    processor_kwargs={"model_max_length": 8192},  # 分词器级别的词元数量上限
)

这个检查点自带它的“配方”:它的查询和文档标记词元、它的投影头、它的打分跳过列表。对于微调,你通常想保留所有这些,只改变你的数据需要的东西。首先要检查的是长度配置,因为许多已发布的检查点把文档限制在 180 到 512 个词元(见为什么要微调?),而我的医学段落长达 1,400 个词元。mLateOn 系列已经支持骨干模型的完整 8192 词元上下文,但如果你的起始检查点带有长度限制,就把它们解除:

# 让模型读取完整文档,而不是它训练时用的截断长度,
# 例如 GTE-ModernColBERT-v1 自带的 query_length=48 和 document_length=300
model[0].query_length = None
model[0].document_length = None

取消了每个任务的长度限制后,截断就会回退到分词器的 model_max_length,这就是为什么我在上面加载时配置了这个限制。

我还做了一个改动,添加了一个标点符号跳过列表,把标点符号词元从文档侧的打分和存储中排除。在一个四路消融实验(无跳过、跳过标点、跳过停用词、两者都跳)中,它在质量上小幅胜出,而且在这个数据上免费让文档索引缩小了 9.6%:

import string

# model[2] 是 MultiVectorMask 模块
model[2].skiplist_words = list(string.punctuation)
model[2].resolve_with_tokenizer(model.tokenizer)  # 词元 ID 会被缓存,所以修改后要重新解析

从基础 Transformer 构建

你也可以让 MultiVectorEncoder 指向任何基础 Transformer,它会自动为你追加一个全新的、随机初始化的词元级投影:

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 生成上下文相关的词元嵌入,一个词元级 Dense 层把每个嵌入投影到 128 维,一个 MultiVectorMask 决定打分时哪些词元算数,以及一个词元级 Normalize 层。投影是随机初始化的,所以这个模型需要训练才能有用。有趣的是,这对强大的密集嵌入骨干模型也适用。在我的实验中,在 Alibaba-NLP/gte-modernbert-base 上加一个全新的投影,仅凭投影和 25k 个训练对,就达到了与现有检查点起点相差 0.03 以内的水平。

经典的 ColBERT 分词技巧([MASK] 查询扩展、[Q] / [D] 前缀词元、文档长度限制、标点符号跳过列表)默认都是关闭的,可以配置。完整的设置列表请参阅创建自定义模型。顺便说一句,我在领域微调中测试了四种配置的 [MASK] 查询扩展,没有一种产生可测量的差异,所以你不必觉得必须使用经典配方。

应该选择哪个起点?

我在准备这篇博客时直接测量了这一点,取了六个起点,用相同的配方在 MIRIAD 的 25k 个医学问答-段落对上训练每个模型,然后在 1,000 个留出的问题上对 50,000 个段落的语料库进行评估:

结果让我很惊讶,而且这个结果在两个模型家族中都得到了验证。*无监督检查点适应新领域的能力远好于它们已经完成训练的“成品”兄弟,尽管起点更低,却能反超。这些检查点位于大规模对比预训练之后、通用检索的监督微调之前,所以它们拥有完整的后期交互结构,却没有那些领域训练需要“撤销”的通用调优。相比之下,成品检查点在我尝试的每个学习率下都几乎不动,甚至出现退化。

所以,如果你喜欢的模型家族发布了预监督检查点,就从那里开始。如果没有,在强大的检索预训练骨干上加一个全新投影是紧随其后的好选择。从完全完成的检查点继续微调是领域适应中最弱的选择,尽管它是最自然感觉的选项。

数据集

MultiVectorEncoderTrainer 使用 datasets.Datasetdatasets.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 个词元)。像这样的简单(查询,相关段落)对是为你自己的领域收集检索训练数据最容易的方式,而且你会看到,它们就是你所需要的全部。

本地数据

你也可以使用 load_dataset 加载常见文件格式的本地数据:

from datasets import load_dataset

dataset = load_dataset("csv", data_files="my_file.csv")
# 或者
dataset = load_dataset("json", data_files="my_file.json")

如果你的本地数据需要预处理,你可以使用 datasets.Dataset.from_dict 用一个字典(键是列名,值是列表)来初始化数据集:

from datasets import Dataset

queries = []
documents = []
# 打开文件,进行预处理、过滤、清洗等操作
# 然后追加到列表中

dataset = Dataset.from_dict({
    "query": queries,
    "document": documents,
})

数据集格式

重要的是,你的数据集格式要与你的损失函数匹配(或者你选择一个与数据集格式匹配的损失函数)。验证数据集格式是否与损失函数兼容需要两个步骤:

  • 如果你的损失函数根据损失函数概览表需要一个标签(Label),那么你的数据集必须有一个名为 "label" 或 "score" 的列。这个列会自动被当作标签。

  • 所有不叫 "label" 或 "score" 的列根据损失函数概览表都被视为输入(Input)。剩余列的数量必须与你选择的损失函数的有效输入数量匹配。这些列的名字无关紧要,只有顺序重要。

除此之外,还有两个多向量特有的约定:

  • 位置式查询和文档分配:第一列被嵌入为查询,所有后续列被嵌入为文档,不管列名是什么。这个默认行为可以通过标准的 router_mapping 训练参数按列覆盖。

  • 知识蒸馏格式:每个候选文档一列,即 (query, document_1, ..., document_N, scores),其中 scores 是每行 N 个教师模型分数的列表。对于把查询和文档 ID 与单独的文本数据集一起存储的 KD 数据集(例如 lightonai/ms-marco-en-bge),你可以使用 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,  # 每个块编码多少个文档:限制内存,不影响质量
)

mini_batch_size 参数通过按这个大小分块编码文档来限制内存,而有效对比批次大小(我在下面的运行中是 128,在我的消融实验中更大的批次没有带来更多收益)仍然可以自由选择。GradCache 保证无论块大小如何,结果都完全相同,所以对于更小的 GPU,只需降低它,代价只是墙钟时间变长。当你的文档长度差异很大时,可以考虑它的兄弟参数 mini_batch_num_tokens,它把每个块打包到总词元预算而不是文档数量,这样一块异常长的文档永远不会让你的内存飙升(我的 mini_batch_size=16,每篇文档大约 940 个词元,相当于 mini_batch_num_tokens=15_000)。

一个多向量特有的陷阱是,对比损失默认 scale=1.0,而密集嵌入的等价损失默认 scale=20.0。那个 20.0 存在是因为余弦相似度是 [-1, 1] 中的单个值,范围太窄,无法形成尖锐的 softmax。而 MaxSim 分数是对每个查询词元的一个最佳匹配相似度求和,所以它已经大致覆盖了 [0, query_length] 的范围:一个 32 词元的查询最高可以得 32 分。所以不要从密集训练脚本中复制 scale=20.0,因为它会让 softmax 饱和,毁掉你的梯度。

对于从更强的教师模型进行蒸馏——这也是最强通用后期交互模型的训练方式——请参阅 MultiVectorDistillKLDivLoss训练概览文档中的知识蒸馏标签页。

训练参数

你可以使用 MultiVectorEncoderTrainingArguments 类来自定义训练过程。这个类让你调整可以影响训练速度并帮助你理解训练过程中发生什么的参数。

关于最有用的训练参数的更多信息,请查看多向量编码器 > 训练概览 > 训练参数。值得一读,以便充分利用你的训练。

下面是一个例子,使用了我实际训练运行中的值:

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,  # 如果你的 GPU 支持 FP16,设为 True
    bf16=True,  # 如果你的 GPU 支持 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",  # 将用于 Trackio、W&B 等
)

其中几个参数值得说明一下:

  • prompts:训练不会自动应用模型中存储的提示词,所以要把它们显式映射到你的训练列上。这里就是检查点的 [Q] 标记对应问题列,[D] 对应段落列,让训练与推理保持一致。

  • max_length(故意不设置):这个参数只在训练期间限制分词,用于当你想要比模型完整服务长度更便宜的训练时。我测量了这个捷径在这个数据上的代价。在 512 个词元下训练损失了大约 0.015 的 NDCG@10,换来大约 2 倍的速度,而且这个差距不会随着更多数据而缩小,因为模型根本看不到被截断的内容。保持不设置,让训练与推理匹配,除非你更需要速度而不是质量。

  • learning_rate=1e-4:在从 5e-6 到 2e-4 的扫描之后,这个比通常更高的学习率效果最好。

评估器

为了在训练期间跟踪模型性能,你可以向训练器传递一个 eval_dataset 来计算评估损失,但具体的检索指标信息量要大得多。Sentence Transformers 为多向量模型内置了以下评估器:

对于领域微调,用你自己的留出数据构建的 MultiVectorInformationRetrievalEvaluator 才是最重要的。构建它的一个技巧是,语料库应该足够难,才能区分不同模型。在我的案例中,MIRIAD 的问题是从它们自己的源段落生成的,这让检索异常容易。仅针对 10k 个黄金段落,几乎所有模型的 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")

# 黄金标准:1,000 个评估问题,每个问题对应一个专属段落,评估集里大约 10,000 个不重复段落作为初始语料库

```python
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"]]}

# 干扰项:训练集中不重复的段落,让“大海捞针”更真实
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)

训练器(Trainer)

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. 加载起始检查点:已经做过对比预训练,但还没经过监督微调
    # 如果内存够用,训练时最好用 fp32 格式加载
    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. 解除每个任务的长度限制,让训练和推理都能看到完整的医学段落
    model[0].query_length = None
    model[0].document_length = None

    # 3. 打分时跳过标点符号:小提升,索引还能缩小 9.6%
    model[2].skiplist_words = list(string.punctuation)
    model[2].resolve_with_tokenizer(model.tokenizer)

    # 4. 加载 100 万个医学问答对
    train_dataset = load_dataset("tomaarsen/miriad-4.4M-split", split="train").select(range(1_000_000))

    # 5. 批内负样本 + GradCache:有效批次大,内存占用可控
    loss = CachedMultiVectorMultipleNegativesRankingLoss(model=model, mini_batch_size=16)

    # 6. 一个轻量开发集评估器,训练时观察进度:500 个留出问题
    #    对应评估集里约 10,000 个不重复段落。完整的 20 万协议之后再做。
    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. 训练参数,如上文讨论
    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,  # 如果你的 GPU 支持 FP16,可以设为 True
        bf16=True,  # 如果你的 GPU 支持 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=run_name,
    )

    # 8. 创建训练器并开始训练
    trainer = MultiVectorEncoderTrainer(
        model=model,
        args=args,
        train_dataset=train_dataset,
        loss=loss,
        evaluator=dev_evaluator,
    )
    trainer.train()

    # 9. 保存训练好的模型
    model.save_pretrained(f"models/{run_name}/final")

    # 10.(可选)推送到 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 上,峰值显存 17.5 GB,运行了 14.5 小时。每一个选择都是经过实测比较胜出的,而不是靠猜。

对于预算更小的读者,我的扩展实验表明,10 万对(75 分钟训练)与完整 100 万对运行的 NDCG@10 差距在 0.012 以内。大部分收益在第一个小时就能获得。

回调(Callbacks)

MultiVectorEncoder 训练器支持各种 transformers.TrainerCallback 子类,包括:

  • WandbCallback:如果安装了 wandb,可将训练指标记录到 W&B

  • TensorBoardCallback:如果可以使用 tensorboard,可将训练指标记录到 TensorBoard

  • CodeCarbonCallback:如果安装了 codecarbon,可追踪训练期间的碳排放

通过 report_to 训练参数启用这些功能,例如 report_to=["wandb", "codecarbon"],并安装相应的依赖。默认值是 "none",而 report_to="all" 会激活所有已安装依赖的集成。

更多关于这些回调以及如何创建自己的回调,请参阅 Transformers 回调文档

多数据集训练

通常,性能最好的通用模型是在多个数据集上同时训练的。然而,由于每个数据集的格式不同,这种方法可能很有挑战性。幸运的是,MultiVectorEncoderTrainer 允许你在不需要统一格式的情况下在多个数据集上训练。此外,它还提供了对每个数据集应用不同损失函数的灵活性。以下是一次训练多个数据集的步骤:

  • 使用 datasets.Dataset 实例的字典(或 datasets.DatasetDict)作为 train_dataset(也可以作为 eval_dataset)。

  • (可选)使用一个损失函数字典,将数据集名称映射到损失函数。仅当你希望对不同数据集使用不同损失函数时才需要。

每个训练/评估批次只会包含来自其中一个数据集的样本。从多个数据集中采样批次的顺序由 MultiDatasetBatchSamplers 枚举定义,可以通过 multi_dataset_batch_sampler 传递给 MultiVectorEncoderTrainingArguments。有效选项包括:

  • MultiDatasetBatchSamplers.ROUND_ROBIN:从每个数据集轮流采样,直到其中一个耗尽。使用这种策略,可能不会用到每个数据集的所有样本,但每个数据集被采样的机会是均等的。

  • MultiDatasetBatchSamplers.PROPORTIONAL(默认):根据每个数据集的大小按比例采样。使用这种策略,每个数据集的所有样本都会被用到,较大的数据集被采样的频率更高。

评估

为了了解微调模型的表现,我在 MIRIAD 评估集上对四个架构家族的 50 多种检索模型配置进行了评估。评估集的构建方式与上面的 评估器 部分完全一致:1,000 个留出的医学问题,在 200,000 个不重复段落中搜索(10,000 个黄金段落隐藏在 190,000 个来自训练集的去重干扰项中)。这个语料库是 应该选择哪个起点? 中 50,000 段语料库的四倍,因此两个表格之间的分数不可比较。

主要结果如下,完整表格在下面的折叠部分:

微调模型位居榜首,比任何架构中最强的零样本模型高出 +0.062 NDCG@10。换句话说,最强的零样本模型在 75.8% 的查询中将正确段落作为第一个结果返回,而微调模型在 84.9% 的情况下做到了这一点,将排名第一的错误率降低了三分之一以上。

架构模式同样清晰,表格顶部完全是晚期交互(late interaction)模型。在长文档上,每个 token 一个向量优于每个文档一个向量,即使在匹配的训练数据和匹配的主干网络下也是如此。DenseOn 和 LateOn 共享训练数据和架构,只是头部不同,晚期交互兄弟模型以 +0.12 的优势胜出,多语言对(mDenseOn 和 mLateOn)以 +0.13 的优势复制了这一结果。规模也无法拯救单向量模型。Qwen3-Embedding-4B 是最强的密集模型,其活跃(非嵌入)参数大约是我的模型的 33 倍,但仍然差了 0.13,而 8B 版本的得分低于 4B。

BM25 的表现也出人意料地好,击败了所有稀疏模型、所有受截断限制的多向量模型,以及除三个密集模型外的所有模型:数十亿参数的 Qwen3-Embedding-4B8B,以及 voyage-4-nano,后者通过读取完整的 32k token 上下文,仅以 0.006 的优势险胜。不过,不要指望这能迁移到你自己的数据上。MIRIAD 的问题是从段落中生成的,因此查询与其黄金段落之间的词汇重叠远大于典型检索场景,而 BM25 的无限上下文长度使其能够利用每一个重叠的词,而大多数神经检查点会截断。BM25 基线成本低廉,总是值得运行,只是不要指望这种优势。

完整领域概览,按分数排序并按架构家族着色。

标记为 @N 的模型是在其文档长度上限提升到 N 个 token 的情况下评估的,因为它们原生的上限(180 到 512 个 token)会截断平均 941 个 token 的段落。对于每个多向量模型,这种提升相对于原始配置带来了 +0.08 到 +0.24 的 NDCG@10 提升,即使是密集的 DenseOn 也从同样的处理中获得了 +0.03 的提升。

请注意,这并不意味着 multi-vector-encoder/mLateOn-medical 在所有领域都是最强的模型。它只是在我的领域中最强。这完全没问题,因为我只需要这个模型在我的数据上表现良好。

不要低估在你的领域微调多向量模型的力量。在单张消费级 GPU 上训练 14.5 小时,就产生了一个在这个数据上没有任何通用检索器能接近的模型,而且配方只是一个脚本,没有教师模型,也没有挖掘负样本!

优化索引

对多向量检索的合理质疑是索引大小,而这个领域接近最坏情况。每个 token 存储一个向量,我的模型每个段落大约需要 878 个向量,因此 200,000 段语料库在 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,排名第一的准确率保持不变,而只保留四分之一的向量(11.2 GB)仍然得分 0.8991。曲线还在继续(我测量到只保留十分之一的向量,仍然有 0.8765),但一旦量化进入考虑范围,就没有太多理由把池化推得那么远,这正是下面虚线所展示的。

虚线是真实部署可能的样子。我给了 Omar Khattab 模型和基准的早期访问权限,他用 fast-plaid 在 1 位残差量化下测量了这些配置,使用紧凑的 17 位质心 ID 和 18 位文档 ID 代替普通的未打包 64 位整数,再加上文档侧剪枝:

第一行比原始嵌入小 13 倍,NDCG@10 只损失 0.0155。这比池化曲线上的任何一点都好得多。量化缩小每个向量,而池化和剪枝减少你保留的向量数量,所以它们可以组合使用,量化是应该首先使用的。进一步推进,最后一行达到 1.45 GB,比 Qwen3-Embedding-8B 的 fp16 嵌入(1.64 GB)还小,同时得分高出 0.0895。多向量索引太大的质疑,在正确配置的索引面前站不住脚。

这里的剪枝是朴素的,只是为了证明 token 减少可以在量化之上工作,所以把下面两行看作下限而不是前沿。如果你根本不想手动调整量化,配套文章的 索引 部分涵盖了 fast-plaid、Qdrant、Weaviate 和 Vespa。

多向量检索的成本只取决于其索引。这个语料库的原始嵌入是 45 GB,而正确配置的索引至少小 7 倍,准确率几乎相同。索引值得你像关注检查点一样关注。

致谢

感谢 Omar Khattab优化索引 中测量了量化和剪枝的索引配置,以及围绕晚期交互索引成本的讨论。

其他资源

训练示例

这些页面包含带解释的训练示例以及训练脚本的链接。你可以用它们来熟悉多向量训练循环:

  • MIRIAD:医学检索的领域特定训练,是这篇博文配方的更早、更简单的版本

  • MS MARCO:对比学习和知识蒸馏配方

  • 多模态:ColPali 风格的视觉文档检索训练

  • PEFT 适配器:使用 LoRA 进行参数高效微调

文档

为了进一步学习,你可能还想探索以下关于 Sentence Transformers 的资源:

这里还有一个你可能感兴趣的高级页面:

以及配套博文,涵盖使用这些模型的所有内容:

本文提到的模型 10

本文提到的数据集 3

更多来自我们博客的文章

社区

· 注册登录 发表评论

本文提到的模型 10

本文提到的数据集 3

术语表

多向量模型
一种嵌入模型,为每个词元生成一个向量,通过 MaxSim 算子计算查询和文档的相似度,能保留细粒度信息。
MaxSim
一种打分机制,每个查询词元找到最匹配的文档词元,并将所有匹配分数求和,用于多向量模型的相似度计算。
词元
文本的最小单位,可以是单词或子词,模型处理文本的基本单元。
NDCG@10
一种衡量检索排序质量的指标,考虑前 10 个结果的排序相关性,值越高越好。
批内负样本
训练时,将同一批次中的其他文档作为当前查询的负样本,增加训练效率。
GradCache
一种梯度缓存技术,允许使用更大的有效批次大小,同时控制内存占用。
预监督检查点
在大型语料上经过对比预训练但未进行监督微调的模型,通常具有更好的领域适应能力。
索引
存储嵌入向量的数据结构,用于快速检索,多向量模型的索引通常较大。
量化
通过降低数值精度(如从 fp16 到 1 位)来减小索引大小,同时尽量保持性能。
池化
通过聚类或平均减少向量数量,如 HierarchicalTokenPooling,以压缩索引。

生产区

使用这篇文章

复制包含 frontmatter、中文正文和英文来源的 Markdown。

我的研究区

记录自己的判断,不会写回原文或公开页面。

一句话说清这篇材料能支撑什么角度。

每行一个,最多 20 个。

支持 Markdown。要核对的数据、可复用的方法、反方观点或补充来源。