交叉编码器重排:两段式检索的精排器


文档摘要

交叉编码器重排:两段式检索的精排器 本节摘要:双塔编码器(Bi-encoder)把查询和文档各自独立嵌入,靠余弦排序——快,但两路编码永不相见,无法区分主题相近却只有一篇真正回答了查询的文档。交叉编码器(Cross-encoder)把两者拼成一个序列一起读,每个文档 token 都能注意到每个查询 token,精度碾压但吞吐骤降:对千万级语料每条查询要千万次前向,不可行。本节实现一个微型交叉编码器(一个 transformer 块 + 多头注意力 + 回归头),搭出「双塔取 top-N → 交叉编码器精排到 top-K」的两段式流水线,在小固定语料上测量「延迟 vs 质量」曲线,并教你根据延迟预算选对 N。

交叉编码器重排:两段式检索的精排器

本节摘要:双塔编码器(Bi-encoder)把查询和文档各自独立嵌入,靠余弦排序——快,但两路编码永不相见,无法区分主题相近却只有一篇真正回答了查询的文档。交叉编码器(Cross-encoder)把两者拼成一个序列一起读,每个文档 token 都能注意到每个查询 token,精度碾压但吞吐骤降:对千万级语料每条查询要千万次前向,不可行。本节实现一个微型交叉编码器(一个 transformer 块 + 多头注意力 + 回归头),搭出「双塔取 top-N → 交叉编码器精排到 top-K」的两段式流水线,在小固定语料上测量「延迟 vs 质量」曲线,并教你根据延迟预算选对 N。读完本节,你能说清为什么交叉编码器总在第二段、N 的拐点为何总在 20~50,以及把 N 设成等于 K 会怎样让提升归零。

对应原课程:Phase 19 · Lesson 66 · reranker-cross-encoder(原英文 phases/19-capstone-projects/66-reranker-cross-encoder/docs/en.md)。本节属第 20 章「毕业项目」的进阶 RAG 赛道。

学习目标

阅读完本节,你应当能够:

  1. 输入形状、参数量、每查询成本三个维度,区分双塔检索器与交叉编码器重排器。
  2. 从零实现一个小型交叉编码器:一个消费打包好的 (query, document) 序列、吐出单个相关性标量的 transformer 块。
  3. 搭出两段式「检索-重排」流水线:廉价检索器取 top-N,交叉编码器把 N 精排到 top-K,返回 K。
  4. 在小固定语料上测量延迟 vs 质量权衡,为给定延迟预算选对 N。

一、问题与直觉

双塔编码器把查询和文档映进同一向量空间,按余弦排序。两个编码互不可见,模型必须把文档的一切有用信息压进单个向量——对查询盲目。这很快(索引期一文一嵌入,查询期一查询一嵌入),也是语料规模上唯一可行的排序方式。

代价是精度。两篇整体主题相同的文档,即便只有一篇真正回答了查询,嵌入也几乎一致——双塔分不开。交叉编码器通过一起读查询与文档来解决这个问题:模型收到 [query] [SEP] [document] 作为单一序列,跨拼接处跑完整注意力,产出一个相关性标量。文档的每个 token 都能注意到查询的每个 token,带着完整上下文打分。

代价是吞吐。双塔嵌入一次,查询无数次;交叉编码器每个 (query, document) 对跑一次。千万级语料每条查询就是千万次前向,请求预算内跑不动。解法是分段:双塔取 top-N,交叉编码器把 N 精排到 top-K。N 很小(50~200),交叉编码器的质量提升集中在要紧处;总延迟留在预算内;总质量是交叉编码器的质量,受限于双塔在 N 处的召回。

交叉编码器的输入形状

标准打包是 [CLS] query_tokens [SEP] document_tokens [SEP]。CLS 位置的输出喂给一个线性头,输出相关性标量。有的实现用均值池化代替 CLS,差异不大。关键是模型每对产出一个数

2200 万参数的交叉编码器(发表版 ms-marco-MiniLM-L-6-v2 量级)是典型生产点;更小的模型质量掉得比省的延迟快。更大的(如 5.68 亿参数的 bge-reranker-v2-m3)留给离线重排或 K 很小的首页重排。

延迟 vs 质量

两段式流水线只有一个可调量:N。在留出查询集上把 N 从 5 扫到 100,就得到曲线。

N 第二段 recall@1 每查询交叉编码器前向数 延迟
5 0.62 5
20 0.81 20
50 0.86 50
100 0.86 100 很高

上表数字是形状示意,非本固定语料实测——但形状是真的:总有一个拐点在 20~50 候选处,重排提升在此饱和;过了拐点就是白花钱。从评估曲线加延迟预算里挑 N。交叉编码器无法把召回抬过双塔在 N 处的召回,所以低 N 同时封顶了质量与延迟。

二、从零实现

code/main.py 实现:

  • CrossEncoder —— 小型 torch.nn.Module:token 嵌入、一个含多头注意力与前馈的 transformer 块、均值池化头产出一个标量。
  • tokenize_pair(query, document) —— 把两串打包成单一 id 序列,带标记边界的 type id,确定性强、纯标准库。
  • train_tiny(pairs) —— 在手标 (query, document, relevance) 三元组列表上跑一轮监督训练,让模型在固定语料上产出合理分数。
  • rerank(query, candidates, top_k) —— 生产接口。
  • pipeline(query, retriever, top_n, top_k) —— 两段式流。
  • 一个 demo main(),按第 65 节的模式加载语料,取 top-N,精排到 top-K,并排打印两份列表,报告各段延迟。

交叉编码器核心骨架

import torch, torch.nn as nn class CrossEncoder(nn.Module): def __init__(self, vocab=1000, dim=64, heads=4): super().__init__() self.embed = nn.Embedding(vocab, dim) self.attn = nn.MultiheadAttention(dim, heads, batch_first=True) self.ffn = nn.Sequential(nn.Linear(dim, dim*4), nn.GELU(), nn.Linear(dim*4, dim)) self.norm1 = nn.LayerNorm(dim) self.norm2 = nn.LayerNorm(dim) self.head = nn.Linear(dim, 1) # 回归头:每对一个相关性标量 def forward(self, ids, type_ids): x = self.embed(ids) + self.embed(type_ids) # type embedding 标记 query/doc 边界 a, _ = self.attn(x, x, x) x = self.norm1(x + a) x = self.norm2(x + self.ffn(x)) pooled = x.mean(dim=1) # 均值池化 return self.head(pooled).squeeze(-1)

两段式流水线

def pipeline(query, retriever, cross, top_n=50, top_k=5): # 第一段:廉价双塔取 top-N candidates = retriever.search(query, k=top_n) # 第二段:交叉编码器对 N 个候选打分,取 top-K scored = [(c, cross(tokenize_pair(query, c.text))) for c in candidates] scored.sort(key=lambda x: -x[1]) return scored[:top_k]

运行:

python3 code/main.py

输出显示双塔的 top-N、交叉编码器的 top-K,以及计时小结。交叉编码器单次更慢,但不对全语料跑;两段总延迟留在请求预算内,却能把双塔排第二、第三的答案提上来。

💡 为什么本节训练一个微型模型:真实交叉编码器是微调过的编码器 transformer,生产里你加载 checkpoint 就跑。本节目的不是训出 SOTA 排序器,而是让你看清模型的形状延迟-质量曲线的形状。所以构建一个带单个 transformer 块、4 头注意力、一个回归头的小 nn.Module,从种子确定性初始化,无需磁盘权重也能复现 demo。

三、框架对比

维度 双塔(Bi-encoder) 交叉编码器(Cross-encoder)
输入形状 query、doc 各自编码 [query][SEP][doc] 拼接编码
注意力 两路互不可见 跨拼接处完整注意力
每查询成本 1 次查询嵌入 N 次前向(N 个候选)
语料规模可行性 唯一可行 不可行(千万级)
典型参数量 110M 22M(MiniLM)~568M(bge-v2-m3)

业界对比:HuggingFace sentence-transformersCrossEncoder 类、Cohere Rerank API、Voyage rerank-2、Jina jina-reranker 都是这一范式,差别只在权重与服务化形态。它们的共识与本节一致:双塔检索、交叉编码器重排、N 在 50200、拐点在 2050

四、demo 会藏住的失败模式

交叉编码器不对称。 rerank(q, d)rerank(d, q) 是不同分数。永远先喂查询。若不小心互换,召回崩塌。

N 太低暴露不了 bug。 若设 N = K,交叉编码器无法重排,只能重新加权——提升看起来是零。N 至少取 K 的三倍。

训练数据漏进评估。 若手标训练对包含了评估查询,重排看起来像魔法。即使在固定语料上也要严格分离 train 与 eval。

生产权重稠密。 2200 万参数的交叉编码器,float32 下是 88MB。承诺 p95 < 100ms 前先规划模型服务器内存。

批处理要紧。 真实交叉编码器把 N 个候选放进一个 batch 跑。本节在 _batch_encode 里用 torch.tensor(...) 构造批处理 id 与 type-id 张量,一次前向。跳过批处理,延迟乘以 N。

五、可复用产物

  • CrossEncoder + pipeline:本节的两段式流水线是第 69 节「端到端 RAG 系统」的第二阶段,接在第 65 节混合检索器之后、答案生成之前。
  • N 的拐点经验值(20~50):作为新语料调参的起点,再以第 68 节评估确认。
  • 延迟-质量曲线模板:把 eval_recall 与计时结合,任何新模型/语料都能复用这张表的形状。

生产模式:

  • 把双塔、交叉编码器、N 一起锁版本——改任何一个都使评估失效。
  • (query, document_id) 哈希缓存重排器输出;同查询对稳定语料重排顺序不变,缓存命中白赚延迟。
  • 记录 rank-1 交叉编码器分数;某查询 top-1 分数低于语料特定阈值即越域命中,作为「我不确定」抛给 LLM。

六、练习

  1. 扫 N 找拐点:把 N 从 5 扫到 50,画出精排输出 recall@1,在本固定语料上找到拐点。
  2. 多训几轮:把交叉编码器训 10 轮而非 1 轮,测量每轮正负样本对的分数间距。
  3. 换池化头:把均值池化换成 CLS-token 头,在本固定语料上对比收敛。
  4. 加第二个头:加一个预测「答案是否在文档里」的二分类头,推理时一头排序、一头阈值。
  5. 接真实双塔:把确定性 mock 双塔换成第 65 节的实现,串起两段,测量 top-K 相对纯双塔的变化。

本节要点回顾

  1. 双塔快但盲目,交叉编码器聪明但慢——两路编码互不可见 vs 跨拼接处完整注意力。
  2. 两段式是工程解:双塔取 top-N,交叉编码器精排到 top-K;N 在 50~200。
  3. 拐点在 20~50:重排提升在此饱和,过了就是白花钱。
  4. 交叉编码器每对一个数:[CLS] q [SEP] d [SEP],CLS/均值池化 + 线性回归头。
  5. 生产点 22M 参数(MiniLM 量级);更小质量掉得比省延迟快。
  6. 五大失败:不对称(先喂查询)、N=K 提升归零、训练漏评估、权重稠密规划内存、跳批处理延迟乘 N。
  7. 工程化:三件套(双塔+交叉+N)一起锁版本,按 (query,doc_id) 缓存,top-1 分数做越域检测。

下一节,我们将进入「查询重写 HyDE」——在查询碰到检索器之前先改写它,用 LLM 写一个假答案文档来弥合查询与语料的 token 鸿沟,让答案真正进入 top-N。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U