交叉编码器重排:两段式检索的精排器 本节摘要:双塔编码器(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 赛道。
阅读完本节,你应当能够:
(query, document) 序列、吐出单个相关性标量的 transformer 块。双塔编码器把查询和文档映进同一向量空间,按余弦排序。两个编码互不可见,模型必须把文档的一切有用信息压进单个向量——对查询盲目。这很快(索引期一文一嵌入,查询期一查询一嵌入),也是语料规模上唯一可行的排序方式。
代价是精度。两篇整体主题相同的文档,即便只有一篇真正回答了查询,嵌入也几乎一致——双塔分不开。交叉编码器通过一起读查询与文档来解决这个问题:模型收到 [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 很小的首页重排。
两段式流水线只有一个可调量: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) —— 两段式流。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-transformers 的 CrossEncoder 类、Cohere Rerank API、Voyage rerank-2、Jina jina-reranker 都是这一范式,差别只在权重与服务化形态。它们的共识与本节一致:双塔检索、交叉编码器重排、N 在 50200、拐点在 2050。
交叉编码器不对称。 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 节混合检索器之后、答案生成之前。eval_recall 与计时结合,任何新模型/语料都能复用这张表的形状。生产模式:
(query, document_id) 哈希缓存重排器输出;同查询对稳定语料重排顺序不变,缓存命中白赚延迟。[CLS] q [SEP] d [SEP],CLS/均值池化 + 线性回归头。下一节,我们将进入「查询重写 HyDE」——在查询碰到检索器之前先改写它,用 LLM 写一个假答案文档来弥合查询与语料的 token 鸿沟,让答案真正进入 top-N。