论文
SparDA:稀疏解耦注意力以实现高效的长上下文 LLM 推理
SparDA: Sparse Decoupled Attention for Efficient Long-Context LLM Inference
摘要
稀疏注意力减少了长上下文 LLM 推理的计算和内存带宽。然而,仍然存在两个关键挑战:(1)KV 缓存容量仍然随着序列长度而增长,并且卸载到 CPU 内存会引入 PCIe 传输瓶颈; (2) 稀疏选择步骤本身保留了 $O(T^2)$ 复杂性,并且可以在长上下文中主导注意力成本。我们提出了 SparDA,一种解耦的稀疏注意力架构,除了查询、键和值之外,还引入了第四个每层投影、预测。 Forecast 预测下一层所需的 KV 块,从而实现将 CPU 到 GPU 预取与当前层执行重叠的前瞻选择。由于 Forecast 与注意力查询分离,因此我们的 GQA 实现对每个 GQA 组使用一个 Forecast 头,与原始多头选择器相比,减少了选择开销。 SparDA 添加了 $<$0.5% 的参数,并通过匹配原始选择器的注意力分布来仅训练 Forecast 投影。在两个稀疏预训练的 8B 模型上,SparDA 匹配或略微提高了准确性,并在稀疏注意力卸载基线上提供高达 1.25$倍$ 预填充加速和 1.7$倍$ 解码加速。通过在单个 GPU 上实现更大的可行批量大小,SparDA 的解码吞吐量比非卸载稀疏基线高出 5.3倍$。我们的源代码可在 https://github.com/NVlabs/SparDA 获取。