论文

Sparton:用融合Triton算子降低稀疏检索训练开销

Sparton: Fast and Memory-Efficient Triton Kernel for Learned Sparse Retrieval

模型推理推理加速

摘要

最先进的学习稀疏检索 (LSR) 模型(例如 Splade)通常采用语言建模 (LM) 头将潜在隐藏状态投影到词汇锚定的 logits 矩阵中。随后,通过逐元素运算(ReLU、Log1P)和序列维度上的最大池化,将该中间矩阵转换为稀疏词汇表示。尽管它很有效,但由于词汇表 (V) 的庞大规模,LM 头造成了巨大的内存瓶颈,在最近的模型中,词汇量的范围可以从 30,000 到超过 250,000 个词元不等。实现该矩阵会产生显着的内存瓶颈,限制模型扩展。由此产生的运算符之间的 I/O 开销进一步限制了吞吐量和运行时性能。在本文中,我们提出了 Sparton,这是一种专为 LSR 模型中的 LM 头量身定制的快速内存高效 Triton 内核。 Sparton 采用融合方法,将平铺矩阵乘法、ReLU、Log1P 和最大值归约集成到单个 GPU 内核中。通过直接在原始 logits 图块上执行早期在线缩减,Sparton 避免了在内存中具体化完整的 logits 矩阵。我们的实验表明,与 PyTorch 基线相比,Sparton 内核单独实现了高达 4.8 倍的加速,并且峰值内存使用量降低了一个数量级。 Sparton 集成到 Splade (|V| ~ 30k) 中,可将批量大小增加 33%,训练速度加快 14%,且不损失检索效果。在多语言骨干网 (|V| ~ 250k) 上,这些增益跃升至 26 倍大的批量大小和 2.5 倍更快的训练速度。