跳转到内容
主站 新闻 控制台

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

· Hugging Face
教程模型卡

Sentence Transformers 是一个 Python 库,用于在各种应用场景中使用和训练嵌入模型与重排模型,例如检索增强生成、语义搜索、语义文本相似度等。其 v6.0 更新引入了第四种模型类型:MultiVectorEncoder,用于 ColBERT 风格的后期交互(late interaction)检索,并为此提供了完整的训练方案。在本文中,我将向你展示如何使用它微调一个多向量模型,使其在你的数据上超越通用检索器。该方法还可以从头训练出性能强劲的新多向量模型。下面的所有内容都可以通过 pip install -U "sentence-transformers[train]" 运行。

微调多向量模型涉及多个组件:模型本身、数据集、损失函数、训练参数、评估器和训练器类。接下来我将逐一介绍这些组件,并通过实际示例说明如何使用它们微调出性能强劲的多向量模型。

最后,在评估部分,我将展示:我使用本文介绍的方法,在单张 RTX 3090 上训练了 14.5 小时得到的 multi-vector-encoder/mLateOn-medical 模型,在我的医疗检索评测中轻松超越了我能找到的所有通用检索模型——无论是稠密、稀疏、词法还是多向量模型。

MIRIAD 上的 NDCG@10 与有效参数量:微调后的 mLateOn-medical 以更小的规模达到顶尖表现,超越了最强的通用模型

如果你对微调稠密嵌入模型、稀疏嵌入模型或重排模型感兴趣,可以阅读我之前发布的嵌入模型的训练与微调、稀疏嵌入模型的训练与微调以及重排模型的训练与微调文章。

本文介绍的是如何训练多向量模型。如果你想了解如何使用多向量模型,包括加载、编码以及在向量数据库中建立索引,请参阅配套文章使用 Sentence Transformers 的多向量(后期交互)嵌入模型。

多向量模型是什么?

稠密嵌入模型会将整段文本压缩成一个向量,两个这样的摘要之间的相似度只需通过一次点积计算。多向量模型(也称为后期交互模型或 ColBERT 风格模型)跳过了这一压缩过程。它会为每个词元保留一个小型向量,并使用 MaxSim 算子计算查询与文档之间的得分:每个查询词元都会找到与其最匹配的文档词元,然后将这些得分相加。词元级匹配能够保留单个向量不得不进行平均压缩的细粒度信号,通常可以带来更强的检索效果,但代价是需要更大的索引。

配套文章多向量嵌入模型详细介绍了架构、编码、评分和索引,因此这里不再展开,直接进入训练部分。

稠密嵌入与多向量后期交互

为什么要微调?

在特定领域上微调多向量模型,可以显著提升其检索性能:网页搜索、法律取证、代码搜索和科学文献综述之间的词汇、查询风格以及相关性定义各不相同。由于查询和文档是逐词元匹配的,多向量模型能够捕捉单向量模型通常会平均掉的细粒度领域信号,即使只有适量的领域内微调数据,也能获得很好的效果。

除此之外,大多数已发布的检索模型都是针对短文本段落配置的。经典的 ColBERT 检查点会将文档截断到 180 或 300 个词元,许多流行的稠密模型则截断到 256 或 512 个词元,因为它们的 MS MARCO 风格训练数据很少超过这一长度。如果你的文档较长,这些模型会在评分前悄无声息地丢弃文档中的大部分内容。在我的医疗评测中,文本段落平均长度为 941 个词元;我测得,仅此截断操作就会导致 NDCG@10 最高下降 0.24,影响远大于不同模型架构之间的差异。训练自己的模型时,你可以根据自己的数据需求配置文档长度。

LightOn 在代码检索领域也遇到了同样的问题:通用模型 LateOn 不够理想,因此他们训练了 LateOn-Code。无论你的领域是医疗、法律、金融,还是公司内部文档,都不会有一个为你专门提供的官方模型。本文将展示如何在数小时内,使用一张消费级 GPU 自行构建适合自己领域的模型。

训练组件

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

  • 模型:要微调的模型,或从头构建的架构。

  • 数据集:用于训练和评估的数据。

  • 损失函数:用于衡量模型性能并指导优化过程的函数。

  • 训练参数(可选):影响训练性能、跟踪和调试的参数。

  • 评估器(可选):用于在训练前、训练期间或训练后评估模型的类。

  • 训练器:整合所有训练组件。

下面来详细了解每个组件。

模型

多向量训练为你提供了多种起始点,而起始点的重要性可能超出你的预期。

微调现有的多向量模型

如果你想进一步微调已有的多向量模型,则完全不必担心架构问题:

from sentence_transformers import MultiVectorEncoder
model = MultiVectorEncoder(
"lightonai/mLateOn-unsupervised",
model_kwargs={"torch_dtype": "float32"},
processor_kwargs={"model_max_length": 8192}, # tokenizer 层面的词元数量限制
)

该检查点包含其