论文
FLASH-MAXSIM:用于后期交互检索的 IO 感知融合内核
FLASH-MAXSIM: IO-Aware Fused Kernels for Late-Interaction Retrieval
摘要
后期交互检索(ColBERT、ColPali)通过 MaxSim 运算符对文档的查询进行评分。标准 PyTorch 实现具体化了完整的查询词元 $\times$ 文档词元相似性张量,只是为了减少它。在 ColPali 规模上,这是管道中最大的单个张量(例如 FP16 中的 21 GB 用于 10K 文档),并且限制了推理时的候选集大小和对比训练期间的批量大小。我们提出了 FLASH-MAXSIM (FM),这是一种 IO 感知的融合 GPU 内核,它可以在不具体化张量的情况下计算相同的 MaxSim 分数,并将相同的原理扩展到向后训练。在 A100 上的 ColPali 规模上,FM 相对于部署的分块基线,将推理峰值内存减少了 1.4-2.6$\times$(相对于未分块的急切执行,减少了 4.9-8.9$\times$),并将 MaxSim 算子训练内存减少了两个数量级,从而能够对较大的常驻候选池和对比批量大小进行精确的重新排序,而普通 autograd 无法适应单个 GPU。该内核是一个直接替代品,精确到其规定的 FP32 累积协议下的浮点计算顺序:nDCG@10 在 BEIR 和 REAL-MM-RAG 上与 FP32 参考最多相差 $5\times10^{-4}$。单独的 INT8 路径以高保真度的准确性换取减半索引存储。代码、基准测试脚本和原始结果:https://github.com/roipony/flash-maxsim