用 Sentence Transformers 训练多向量嵌入模型:从单向量检索走向 Late Interaction

2026-08-26 43 预计阅读时间: 1 分钟
来源: huggingface.co AI 摘要 Original link

Disclaimer: This article is an AI-assisted summary. Read it together with the original source when precision matters. The summary may omit context, version differences, or edge cases and is not official documentation.

预计阅读时间:9 分钟

传统文本检索通常把一段查询和一篇文档各压缩成一个向量,再用余弦相似度或点积完成召回。这种方式索引轻巧、吞吐稳定,但也会丢掉词级匹配信号。多向量嵌入模型让一个文本输出一组向量,并在检索阶段执行更细粒度的匹配,因此适合术语密集、实体较多或长文档检索场景。

Sentence Transformers 的训练生态为这一类模型提供了熟悉的基础设施:Transformer 编码器、数据集、训练循环、评估器以及微调工作流。不过,多向量模型的关键不只是“把池化层去掉”,还包括相似度函数、负样本构造和索引成本的重新设计。

单向量与多向量的分界

单向量模型常见的打分形式如下:

[ score(q, d) = q^T d ]

其中查询 q 和文档 d 都只有一个固定维度的表示。多向量模型则保留 token 级或片段级向量:

[ score(Q, D) = \sum_{i \in Q} \max_{j \in D} Q_i^T D_j ]

这个 Late Interaction 思路意味着:查询中的每个向量都在文档向量集合中寻找最相关的匹配,再把各查询位置的最佳匹配累加。

它带来两个直接结果:

  • 查询中的关键术语可以各自匹配文档中的不同位置,不必共同挤进一个全局向量。
  • 训练目标必须约束“正确文档整体得分高于负文档”,而不仅是让两个 pooled embedding 靠近。

代价也很明确:文档索引会保存多个向量,在线打分比一次点积更贵。模型效果、存储成本和延迟必须一起评估。

微调数据仍然决定上限

多向量结构不会自动修复训练数据中的语义偏差。一个实用训练集至少应包含:

  • query:用户会实际输入的查询,避免只使用标题式短语。
  • positive:确实能够回答查询的段落或文档。
  • negative:主题相近但不回答问题的文本,尤其是 hard negative。

对于检索任务,三元组数据通常比孤立的句对更直接:(query, positive, negative)。如果已有点击日志,可以把同一次搜索曝光但未被点击、或后续被快速返回的文档作为候选负样本;但应过滤位置偏差、重复文档和明显的脏数据。

训练和评估时还应固定文档切分策略。若训练使用 180 token 的段落、生产环境却索引整篇长文,模型看到的匹配粒度会发生变化,离线指标通常会失真。

一个可运行的 Late Interaction 最小训练示例

下面的脚本用于理解多向量训练的核心机制,不是 Sentence Transformers 的特定公开 API 调用。它使用 PyTorch 构建一个小型 token 编码器,为每个 token 输出一个向量,并以正负文档的排序损失进行训练。

运行前安装依赖:

pip install torch
python train_multivector.py

将以下内容保存为 train_multivector.py。实际接入 Sentence Transformers 时,可以将这里的 TokenEncoder 替换为已有 Transformer backbone 的 token embedding 输出,并保留相似度和损失设计。

import torch
from torch import nn

# 0 是 padding token;示例词表仅用于让脚本可独立运行。
VOCAB_SIZE = 128
PAD_ID = 0
EMBEDDING_DIM = 64


def late_interaction_score(query_vectors, document_vectors, query_mask, document_mask):
    """Return one score per query-document pair."""
    similarities = torch.matmul(query_vectors, document_vectors.transpose(1, 2))
    similarities = similarities.masked_fill(
        ~document_mask.unsqueeze(1), float("-inf")
    )
    best_document_match = similarities.max(dim=2).values
    best_document_match = best_document_match.masked_fill(
        ~query_mask, 0.0
    )
    return best_document_match.sum(dim=1)


class TokenEncoder(nn.Module):
    def __init__(self):
        super().__init__()
        self.embedding = nn.Embedding(VOCAB_SIZE, EMBEDDING_DIM, padding_idx=PAD_ID)
        self.projection = nn.Linear(EMBEDDING_DIM, EMBEDDING_DIM, bias=False)

    def forward(self, token_ids):
        vectors = self.projection(self.embedding(token_ids))
        return nn.functional.normalize(vectors, p=2, dim=-1)


# 每行分别是 query、正样本文档、负样本文档;真实项目应替换为 tokenizer 输出。
query_ids = torch.tensor([[4, 17, 23, 0], [8, 31, 0, 0]])
positive_ids = torch.tensor([[17, 50, 23, 9, 0], [31, 72, 8, 0, 0]])
negative_ids = torch.tensor([[90, 91, 92, 0, 0], [60, 61, 62, 0, 0]])

query_mask = query_ids.ne(PAD_ID)
positive_mask = positive_ids.ne(PAD_ID)
negative_mask = negative_ids.ne(PAD_ID)

model = TokenEncoder()
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3)
margin = 0.2

for epoch in range(100):
    query_vectors = model(query_ids)
    positive_vectors = model(positive_ids)
    negative_vectors = model(negative_ids)

    positive_score = late_interaction_score(
        query_vectors, positive_vectors, query_mask, positive_mask
    )
    negative_score = late_interaction_score(
        query_vectors, negative_vectors, query_mask, negative_mask
    )

    # 希望正例分数至少比负例高 margin。
    loss = torch.relu(margin - positive_score + negative_score).mean()

    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

    if epoch % 20 == 0:
        print(
            f"epoch={epoch:03d} loss={loss.item():.4f} "
            f"positive={positive_score.mean().item():.3f} "
            f"negative={negative_score.mean().item():.3f}"
        )

这个示例中最值得保留的不是玩具词表,而是三个约束:文档输出多个向量、每个查询向量对文档取最大匹配、损失直接比较正负文档得分。替换为真实模型后,还需要使用 attention mask 排除 padding,并明确是否移除特殊 token。

生产落地时优先测量什么

多向量检索不应只看 Recall@k。建议将评估拆成三个层面:

  1. 检索质量:在固定查询集上计算 MRR、nDCG、Recall@k,并按短查询、实体查询、长尾查询分别查看。
  2. 索引规模:记录每篇文档的平均向量数、向量维度、量化后的磁盘占用与内存占用。
  3. 服务延迟:分别测量查询编码、候选召回、Late Interaction 重排,以及端到端 P95/P99 延迟。

可以先将多向量模型部署为重排器:用现有 BM25 或单向量模型召回前 100 到 500 个候选,再对候选执行精细打分。这样能较低风险验证质量提升,避免一开始就承担全量多向量 ANN 索引的复杂性。

采用建议

当业务问题依赖精确术语、产品型号、代码标识符、医疗或法律实体,且单向量模型经常把“主题相近”误判为“答案相关”时,多向量模型值得进入实验队列。相反,若语料很短、召回量极大且延迟预算极紧,单向量模型通常仍是更合适的基线。

开始微调前,检查这几项:训练数据是否包含高质量 hard negative;文档切分是否与线上一致;评估是否覆盖延迟和索引成本;以及是否先用候选重排验证收益。把这些边界定义清楚,多向量嵌入才能从离线指标走到可维护的检索系统。


相关推荐