跳到主要内容

数据截至 (上游 commit 5cc889d47547)

04 · CrossEncoder 对照

这一章讲什么: 双塔之外的另一半——CrossEncoder 为什么更准、为什么也更慢,它的 predict/rank 怎么用,训练它的损失族,以及它和双塔在真实检索系统里怎么分工。


1. 它要解决的小问题

双塔把 query 和 document 各自独立编码成一个向量,再算余弦——两个文本在向量相见之前从未「见过面」。这带来一个精度天花板:模型无法捕捉词级别的细粒度交互(比如「not」出现在哪一边)。

CrossEncoder 的选择相反:把两个文本拼成一个序列,过一次完整的 transformer,让每一层的交叉注意力都参与,最后由分类头吐一个相关度分数。它回答「这两段文字有多相关」,而不是「各自是什么意思」。

代价写在性质里:

  • 不能预计算:没有「文档向量」这种东西,分数必须逐对现算。
  • 成本随对数线性:N 个候选文档就是 N 次前向,全库扫描不可行。

所以两者不是竞争关系,是流水线关系:

query ──► 双塔 SentenceTransformer ──► 向量近邻搜索 ──► top-100 候选


CrossEncoder 逐对重打分 ──► 精排 top-10

2. 模型形态:同一副骨架,不同的链

CrossEncoder(sentence_transformers/cross_encoder/model.py:35)同样继承 BaseModel——nn.Sequential 模块链的玩法不变,变的是链上装什么。加载一个裸 HF 模型时的默认链(_load_default_modules,cross_encoder/model.py:212-281)按骨干架构分两路:

骨干架构默认链分数来源
普通 encoder(BERT 等)Transformer(task="sequence-classification")分类头 logits,num_labels=1 即回归分数
CausalLM(Qwen/Llama 类)Transformer(task="text-generation") + LogitScore最后一个 token 处 "yes""no" 的 logit 差

LogitScore(cross_encoder/modules/logit_score.py:8)把生成模型变成打分器的办法很直白:取最后位置 logits,算 logit[yes] − logit[no] 的 log-odds(logit_score.py:42-54)。左 padding 由 Transformer 强制,所以最后位置必是真 token(注释,:43-45)。默认链要求 tokenizer 有 "yes"/"no" 两个词,否则直接报错让你自建模块(cross_encoder/model.py:256-264)。

与双塔对照着记:

SentenceTransformerCrossEncoder
输入单文本文本
输出向量(sentence_embedding)分数(scores,shape (batch, num_labels))
默认骨干任务feature-extractionsequence-classification
典型模块链Transformer→Pooling→NormalizeTransformer(+LogitScore)
默认损失CoSENTLossnum_labels==1 ? BinaryCrossEntropyLoss : CrossEntropyLoss(cross_encoder/trainer.py:172-178)

3. 用起来:predict 与 rank

predict(cross_encoder/model.py:568 附近的实现,签名在 :473-568 的多重重载之后)接收句对列表,返回分数数组;activation_fn 默认 num_labels=1 时 Sigmoid、否则 Identity(docstring,model.py:114-116)。

rank(cross_encoder/model.py:752)是 rerank 场景的便利封装:给一个 query + 一组 documents,内部拼成句对批量 predict,返回按分数降序的 [{corpus_id, score, text?}]。文档字符串里的例子(柏林人口 query)显示分数是 logit 尺度——正例 ≈8.4、负例 ≈−4.3(model.py:126-141)。

训练侧与双塔的差异集中在 collator:CrossEncoderDataCollator 不做 tokenize(cross_encoder/data_collator.py:12-23),只把原始文本列装进 batch 并解析出单一的 prompt/task——因为句对要合在一起编码,预处理发生在损失函数内部(各 loss 自己调 model.preprocess 组对)。这也解释了它的两条额外校验:prompt 和 router_mapping 不允许按列配置(data_collator.py:25-45)——列在 CrossEncoder 里是要被拼起来的,不存在「每列一个 prompt」。


4. 损失族:从点对到排序

CrossEncoder 的损失(sentence_transformers/cross_encoder/losses/)按「标签粒度」分层:

损失数据形态干什么文件
BinaryCrossEntropyLoss(句对, 0/1 或 [0,1] 分数)逐对拟合标签,MS MARCO reranker 的经典配方losses/binary_cross_entropy.py:9
CrossEntropyLoss(句对, 类别 id)多分类(NLI 三分类等)losses/cross_entropy.py
MarginMSELoss(query, 正, 负) + 教师分差蒸馏:让学生的分差逼近教师(强双塔/强 CE)的分差losses/margin_mse.py:10
LambdaLoss / RankNetLoss / ListNetLoss / PListMLELoss(query, 文档列表, 标签列表)listwise 排序:直接优化 NDCG 类指标的可微近似losses/lambda_loss.py:103rank_net.py:11list_net.py:10plist_mle.py:45
MultipleNegativesRankingLoss(CE 版)(query, 正, 负…)与双塔同名思路,但分数来自联合编码losses/multiple_negatives_ranking.py

两个值得记住的联系:

  1. MarginMSE 是双塔与 CrossEncoder 之间的桥。 经典的「双塔蒸馏 CE」配方:用 CrossEncoder 教师给 (query, 正/负) 对打分,双塔学生学习分差——mine_hard_negatives(..., output_scores=True) 产出的分数列就是为这类损失准备的(见第 3 章 §4.4)。MS MARCO 示例 train_bi_encoder_margin_mse.py 走的就是这条路(examples/sentence_transformer/training/ms_marco/)。
  2. 数据格式与损失一一对应,mine_hard_negativesoutput_format="labeled-pair" → BCE、"labeled-list" → LambdaLoss(第 3 章 §4.4 的表),造数工具与训练损失是配套设计的。

5. 原理演示:联合编码 vs 独立编码

# 示意,非源码
# 双塔:两个文本从不相见
q_vec = bi_encoder.encode(["柏林有多少人?"]) # (1, 768)
d_vec = bi_encoder.encode(["柏林 2019 年人口 352 万。"])
score = cosine(q_vec, d_vec) # 向量相见即打分

# CrossEncoder:拼成一个序列过一次模型
pair = ("柏林有多少人?", "柏林 2019 年人口 352 万。")
score = cross_encoder.predict([pair]) # 交叉注意力全程参与

重点看:双塔的分数是两个向量的函数,CrossEncoder 的分数是原始 token 序列的函数——后者表达力强得多,前者才能把 N 次前向摊成「全库一次离线编码 + 一次查表」。


6. 关键细节 / 坑

  • num_labels=1 的分数是 logit,不是概率。 没套 Sigmoid 时分数可正可负(如 §3 的 8.4/−4.3),跨模型比绝对值没有意义;predict 的默认激活只在 num_labels=1 时套 Sigmoid。
  • CrossEncoder 不能拿来建索引。 没有向量产出,encode 不存在;把它当双塔用是新手最常见的误用。
  • rerank 的候选数是有讲究的。 太少召回不足、太多延迟爆炸;库里 RerankingEvaluator(sentence_transformer/sentence_transformer/evaluation/reranking.py:27)与 MS MARCO 示例用的量级是 top-100 上下(inferred:常见配方,非代码所述)。
  • 训练数据必须带负例。 全 0/1 标签的句对可以靠 mine_hard_negatives(output_format="labeled-pair") 直接造(binary_cross_entropy.py 的 Recommendations 也是这么建议的)。
  • CausalLM reranker 依赖 tokenizer 有 yes/no 单 token,没有就得自己拼模块(cross_encoder/model.py:256-268)。

7. 代码地图(本章)

主题文件路径符号名
模型主体sentence_transformers/cross_encoder/model.pyCrossEncoderpredictrank
默认链分派sentence_transformers/cross_encoder/model.py_load_default_modules
生成式 reranker 打分sentence_transformers/cross_encoder/modules/logit_score.pyLogitScore
训练器sentence_transformers/cross_encoder/trainer.pyCrossEncoderTrainer.get_default_loss
数据整理(不 tokenize)sentence_transformers/cross_encoder/data_collator.pyCrossEncoderDataCollator
逐对损失sentence_transformers/cross_encoder/losses/binary_cross_entropy.pyBinaryCrossEntropyLoss
蒸馏分差sentence_transformers/cross_encoder/losses/margin_mse.pyMarginMSELoss
listwise 排序sentence_transformers/cross_encoder/losses/lambda_loss.pyLambdaLossRankNetLossPListMLELoss
双塔蒸馏 CE 示例examples/sentence_transformer/training/ms_marco/train_bi_encoder_margin_mse.py全文