Skip to main content

Cross-Encoder

1. 什么是 Cross-Encoder?

Cross-Encoder 是一种交叉编码 结构,输入是“成对的文本”(如查询+文档),模型把两段文本拼接后一起送入同一个 Transformer,直接输出两者的相关性分数。与 Bi-Encoder(双编码器、分别编码后再算相似度)相比,它的精度通常更高,但推理速度更慢。

2. 适用场景与优势

  • 重排序/精排 :在已有候选文档的情况下,用更高精度的 Cross-Encoder 进行排序(Top-K 精排)。
  • 精度优先 :交叉注意力能捕捉更细粒度的交互。
  • 无需向量索引 :直接输出相关性分数,架构简单。
  • 缺点 :无法预编码,推理成本高,不适合超大规模召回。

3. 学习前置知识补充

  • Transformer 基础 :输入格式常见为 [CLS] 文本A [SEP] 文本B [SEP],输出的 [CLS] 向量用于分类或打分。
  • [CLS][SEP] :BERT 类模型的特殊符号,用于聚合序列信息和分隔句子。
  • 交叉注意力 :两个文本被放在同一序列,注意力可以跨文本交互,比各自独立编码更细粒度。
  • 与 Bi-Encoder 的差异 :Bi-Encoder 支持预编码文档、速度快;Cross-Encoder 精度高但需逐对计算。

4. 核心流程与直观理解

1) 将“查询 + 文档”拼接为一个输入序列。
2) 送入同一个 Transformer,进行多层自注意力计算。
3) 取 [CLS] 位置的输出向量,接一个线性层或 MLP,输出相关性分数(如 0~1)。
4) 对多个候选文档,重复上述步骤并按分数排序。

5. Cross-Encoder 与 Bi-Encoder 对比(简表)

维度Cross-EncoderBi-Encoder
编码方式拼接后一起编码各自编码再算相似度
推理速度慢(逐对计算)快(文档可预编码)
精度略低
适用场景精排、Top-K大规模召回
索引需求不需要向量索引需要向量索引/检索

6. 工作步骤(精排场景)

1) 上游召回(常用 Bi-Encoder 或 BM25)得到候选文档集合。
2) 将“查询+候选文档”逐对输入 Cross-Encoder,得到相关性分数。
3) 按分数排序,取前 K 条作为最终结果。

7. 纯 Python 可运行示例(用简单打分函数模拟 Cross-Encoder)

说明:为避免依赖大模型,这里用“关键词重叠 + 位置权重”模拟一个简单的打分器,方便本地直接运行、观察流程。

# 导入标准库
from typing import List, Tuple
import math

# 简单的“交叉打分器”模拟(非真实模型,仅演示流程)
def simple_cross_encoder_score(query: str, doc: str) -> float:
# 将文本按空格切分为词,实际模型会用分词器并嵌入
q_tokens = query.lower().split()
d_tokens = doc.lower().split()
score = 0.0
# 遍历查询词,计算在文档中的匹配情况
for qi, qtok in enumerate(q_tokens, 1):
for di, dtok in enumerate(d_tokens, 1):
if qtok == dtok:
# 使用位置的倒数作为权重,模拟“越靠前越重要”
score += 1 / (qi + di)
# 使用 log 平滑,避免长文本过大
return math.log(1 + score)

# 对一组候选文档进行精排
def rerank(query: str, docs: List[str]) -> List[Tuple[str, float]]:
# 计算每个文档的分数
scored = [(doc, simple_cross_encoder_score(query, doc)) for doc in docs]
# 按分数降序排序
scored.sort(key=lambda x: x[1], reverse=True)
return scored

# 示例数据
query = "python 面向对象 基础"
docs = [
"python 基础 语法 入门 教程",
"java 面向对象 设计",
"python 面向对象 进阶 实战",
"python 数据分析 numpy 入门"
]

# 运行精排
ranked = rerank(query, docs)

# 打印结果
print("查询:", query)
print("候选文档:")
for d in docs:
print("-", d)
print("\n精排结果:")
for i, (doc, score) in enumerate(ranked, 1):
print(f"{i}. 分数={score:.4f} | 文档={doc}")

8. 可选:使用 sentence-transformers 的真实 Cross-Encoder

需要安装 sentence-transformerstorch,若网络受限可跳过。本示例展示真实用法,代码可独立运行(需成功安装依赖)。

# (可选)先安装依赖:pip install -U sentence-transformers torch
# 导入库
from sentence_transformers import CrossEncoder
import numpy as np

# 加载一个轻量模型(体积较小,适合演示)
model_name = "cross-encoder/ms-marco-MiniLM-L-6-v2"
model = CrossEncoder(model_name)

# 构造查询-文档对
pairs = [
("python 面向对象 基础", "python 面向对象 进阶 实战"),
("python 面向对象 基础", "java 面向对象 设计"),
("python 面向对象 基础", "python 数据分析 numpy 入门")
]

# 预测分数
scores = model.predict(pairs)

# 排序
indices = np.argsort(scores)[::-1]
print("按分数降序:")
for i in indices:
print(f"分数={scores[i]:.4f} | 对={pairs[i]}")

9. 训练数据与常见损失

  • 数据格式 :常见为 (query, positive_doc, negative_doc) 或 (query, doc, label)。
  • 二分类损失 :相关/不相关。
  • 三元组/对比损失 :拉近正样本分数,拉远负样本分数。
  • Listwise 损失 :按列表整体排序优化(初学者可先用二分类或三元组)。

10. 实践建议

  • 用在精排 :先用 Bi-Encoder/BM25 召回,再用 Cross-Encoder 精排。
  • 控制候选集大小 :Top-K 通常 50~1000,越大越慢。
  • 批处理推理 :合并多个 pair 一起推理,充分利用 GPU/CPU。
  • 蒸馏思路 :可用 Cross-Encoder 生成标签,再训练轻量 Bi-Encoder,加速部署。
  • 缓存热门查询 :对高频查询缓存精排结果,减少重复计算。

11. 常见问题

1) 为什么不能预编码?
因为两个文本必须一起过同一个编码器,模型需要交叉注意力。
2) 如果候选很多怎么办?
先用便宜的召回(Bi-Encoder/BM25)缩小集合,再精排。
3) k 取多少合适?
视场景而定,常见 50~1000;越大越耗时。
4) 需要归一化分数吗?
同一模型输出的分数可直接排序;跨模型比较时可做 min-max 归一化。

12. 总结

  • Cross-Encoder 精度高、适合精排,但推理成本大、无法预编码。
  • 建议与 Bi-Encoder 搭配使用:Bi-Encoder 召回,Cross-Encoder 精排。
  • 入门可先用模拟打分函数理解流程,再尝试 sentence-transformers 的真实模型。
  • 控制候选规模、批量推理、缓存热门查询,是落地时的常用优化手段。