从双塔走到交叉:AI交叉编码器笔记
前阵子AI交叉编码器笔记连续改了几轮,趁还记得写成备忘。
问题是什么
在做文本相似度排序时,一开始选了双塔模型。原因简单:它能把每个文本编码成固定长度的向量,之后做相似度计算只需要一个余弦相似度或点积,候选集里一百万个文档也能很快筛出来。
但实际用下来发现几个问题。
双塔模型在粗排阶段表现还可以,一旦到了精排阶段,文本之间的细粒度交互信息丢得太多。比如两句话字面上差不多,但意思完全相反,双塔模型编码出来的向量可能非常接近,排序时就把不相关的结果推到了前面。这在问答匹配、重复文本检测这类场景里尤其明显。
另一个问题是对短文本的处理。短文本本身信息量少,单独编码时很难捕捉到上下文的细微差异,但长文本单独编码又会让向量维度膨胀,计算成本上去后效果却不一定跟着涨。
试过一些补救:调大向量维度、加额外的特征工程、在训练时用更难的负样本,但提升都有限。慢慢意识到,双塔模型的架构本身就有上限,再怎么调也突破不了这个天花板。
为什么考虑交叉编码器
交叉编码器的思路完全不同。它不单独编码两个文本,而是把两个文本拼接起来,作为一个完整输入喂给模型。模型内部能看到两个文本之间的所有交互信息,包括词对齐、语义冲突、上下文依赖这些双塔模型看不到的东西。
在 Hugging Face 上用现成的 Cross Encoder 试了一下,准确率确实比双塔高一个档次。对一些之前排错的结果,交叉编码器能正确判断出不相关,把更相关的文本推到前面。
但问题也很明显:推理速度慢。双塔模型可以先把所有候选文本编码成向量存起来,查询时只需要编码一个查询文本,然后用余弦相似度筛选。交叉编码器每次都要把查询文本和每个候选文本拼接起来走一遍推理,候选集太大时这个开销扛不住。
所以要做两件事:一是把交叉编码器用起来,二是想办法让它跑得够快。
先做最简单的版本
第一版方案很简单:小候选集场景直接上 Cross Encoder,大候选集场景用双塔做粗排,取 top-k 再用 Cross Encoder 做精排。
在 Python 里用 sentence-transformers 跑起来很快:
from sentence_transformers import CrossEncoder
# 加载预训练模型
model = CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')
# 查询和候选文本
query = "如何用 Python 读取 CSV 文件"
candidates = [
"Python 读取 CSV 的三种方法",
"Python 写入 CSV 文件的完整指南",
"Java 读取 CSV 文件的最佳实践"
]
# 打包成对,送入模型
pairs = [[query, cand] for cand in candidates]
scores = model.predict(pairs)
# 按分数排序
ranked = sorted(zip(candidates, scores), key=lambda x: -x[1])
这个方案在候选集小于一千的场景下跑得还行,平均响应时间能控制在 200ms 以内。但在候选集超过一万时,延迟直接飙升到几秒,这种水平在线上服务里基本不可用。
踩坑记录
显存不够
第一版跑起来后发现显存占用比预期高。交叉编码器推理时需要一次性处理多个输入对,批量大小设得太小,GPU 利用率上不去;设得太大,显存又不够用。
后来查了一下 sentence-transformers 的源码,发现它的默认批量大小是 32,在 8GB 显存的机器上很容易 OOM。手动调到 16 后稳住了,但推理速度又慢了一截。
这里学到一个经验:显存不够时不要只想着减批量大小,可以换更小的模型或者开启混合精度训练。sentence-transformers 支持 FP16 推理,改一下就能把显存占用砍掉将近一半:
model = CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')
model.predict(pairs, batch_size=16, convert_to_tensor=True, show_progress_bar=False)
结果不稳定
在同一个输入上跑多次,排序结果有时会不一样。查了一圈发现是模型内部用了 dropout,推理时没有关掉。在 Cross Encoder 里禁用 dropout 后就稳定了:
model.model.eval()
另一个原因是 batch norm 在推理时没有用 running mean 和 variance。这些细节在推理代码里都要显式处理,不然结果就会有随机波动。
端到端部署麻烦
从本地脚本到线上服务之间还有一堆坑。模型加载很慢,冷启动时第一个请求要等好几秒;推理并发处理不当容易挤爆资源;模型文件太大,每次部署都要传几百 MB。
后面想了几个办法:一是用模型量化,把 FP32 模型量化成 INT8,模型体积能缩小 4 倍,推理速度也能提升不少;二是做模型预加载,服务启动时就把模型载入内存;三是实现请求批处理,多个查询可以合并一起推理。
量化用的是 ONNX Runtime,先把模型导出成 ONNX 格式,再量化:
from sentence_transformers import CrossEncoder
import torch
from onnxruntime.quantization import quantize_dynamic, QuantType
# 加载模型
model = CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')
# 导出 ONNX(假设已经实现了导出逻辑)
model.save_pretrained('onnx_model')
# 动态量化
quantize_dynamic(
'onnx_model/model.onnx',
'onnx_model/model_quant.onnx',
weight_type=QuantType.QUInt8
)
量化后模型体积从 400MB 降到 100MB 左右,显存占用也明显降低。在精度上的损失很小,在排序准确率上几乎看不出区别。
架构演进
最开始的架构是双塔单一方案,后面加上了 Cross Encoder 的精排层,再后面又做了缓存和批处理优化。
整个演进过程可以用一个简单的流程图来说明:
这个图的意思是:所有请求先走双塔模型做召回,根据候选集大小决定是否需要做额外粗排。最终都通过 Cross Encoder 做精排,返回排序结果。
除了这个主流程,还加了一些旁路:热点查询的排序结果会缓存起来,相同或相似查询可以直接从缓存读;高频候选的编码向量也会缓存,避免重复计算。
实际效果
用真实数据测了一轮,主要关注三个指标:排序准确率、响应延迟、资源占用。
排序准确率用 NDCG@10 衡量,双塔模型的得分在 0.72 左右,加上 Cross Encoder 精排后升到了 0.79。这个提升在高精匹配场景里特别明显,比如问答匹配、意图识别这些任务。
响应延迟分三种场景看:
- 小候选集(< 100):平均 50ms,大部分时间花在模型推理上
- 中候选集(100 - 1000):平均 150ms,批量推理的收益开始显现
- 大候选集(> 1000):平均 500ms,双塔粗排减少了需要精排的候选数量
资源占用方面,单个服务实例需要 2GB 显存和 8GB 内存。通过批处理和缓存,单实例能扛住每秒 50 个左右的请求。如果需要更高吞吐量,可以横向扩容。
还有哪些边界
这个方案不是万能的,也有一些局限。
对实时性要求特别高的场景,Cross Encoder 的推理延迟仍然偏高。如果需要把延迟压到 50ms 以下,可能需要考虑更轻量的模型或者专门做蒸馏。
跨语言场景下,用单语言 Cross Encoder 效果不够好。需要换多语言模型,但这会进一步增加推理成本。
另一个边界是长文本。Cross Encoder 对输入长度有限制,太长的文本会被截断,导致信息丢失。这个可以通过分块处理或者用长文本模型来解决,但复杂度会上去。
小结
从双塔到交叉编码器的折腾过程,本质上是在效果和效率之间找平衡。双塔模型快但精度有限,交叉编码器准但不够快,最终通过两阶段的架构把两者的优势结合起来。
这次迁移之后,排序质量有明显提升,但技术复杂度也跟着涨了不少。模型量化、缓存策略、批处理、并发控制这些细节都需要仔细调,否则很容易在某个环节掉链子。
后来想了想,其实大多数技术选型都是这样的:没有完美的方案,只有在特定约束条件下的最优解。这次做得对的事情,可能是承认了双塔模型的局限性,愿意花成本去换一个更准但更慢的方案;把事情做对的部分,大概是在效果提升的同时,通过各种优化手段把性能开销控制在一个可接受的范围。
技术项目里,“做得对"和"把事做对"经常打架,这次算是把两个方向都往前推了一小步。
版权声明: 本文首发于 指尖魔法屋-从双塔走到交叉:AI交叉编码器笔记(https://blog.thinkmoon.cn/post/368-ai-cross-encoder-dual-tower-cross-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。