注意力状态与递归注意力¶
FlashInfer 引入了注意力状态的概念,它完全描述了查询与一组键/值对之间的注意力。我们进一步定义了注意力状态上的一个合并算子。这个合并算子通过允许注意力状态的递归合并来促进完整注意力的计算。
假设我们定义 \(s_i = \mathbf{q}\mathbf{k}_i^T\) 作为查询 \(\mathbf{q}\) 和键 \(\mathbf{k}_i\) 之间的 softmax 前的注意力分数。在索引 \(i\) 上的自注意力分数可以推广到索引集 \(I\)
我们也可以将索引 \(i\) 上的值推广到索引集 \(I\)
softmax 函数限制在索引集 \(I\) 内。请注意,\(\mathbf{v}(\{1,2,\cdots, n\})\) 是整个序列的自注意力输出。索引集 \(I\) 的注意力状态可以定义为元组 \((s(I), \mathbf{v}(I))\),然后我们可以定义两个注意力状态的二元合并算子 \(\oplus\) 为(在实践中,我们将 s 与最大值相减以保证数值稳定性,这里为了简单起见省略了它们)
合并算子可以推广到任意数量的注意力状态输入
上述 n 元合并算子与二元合并算子一致,我们可以证明该算子是交换律和结合律的。可以通过合并索引子集的注意力状态来获得整个序列的注意力状态,最终结果在数学上是等效的
注意
广义分数 \(s\) 也被称为对数和指数 (简写为 lse)。
应用¶
请注意,\(\oplus\) 算子是交换律和结合律的,这意味着我们可以安全地将 KV 子集上的自注意力计算卸载到不同的设备上,并以任何顺序合并结果。
到目前为止,FlashInfer 中存在这种递归自注意力形式的几个有趣的应用
- 共享前缀批量解码
许多 LLM 应用程序涉及使用共享长提示的批量解码,FlashInfer 将整个 KV-Cache 上的注意力分解为共享前缀注意力和唯一的后缀注意力。这种分解能够将这些组件卸载到不同的内核实现,从而在长上下文和大型批量大小的场景中实现高达 30 倍的加速。这种分解在长上下文设置中将算子加速 30 倍。请查看 我们的博客文章,了解有关此应用程序的更多详细信息,以及 Cascade Attention,了解如何在 FlashInfer 中使用此功能。
- KV 序列并行性
对于长上下文 LLM 推理/服务,GPU 的批量大小和每 GPU 的头数受到 GPU 内存的限制,默认的并行策略无法使用 GPU 中的所有 SM,从而导致次优性能。受到 GEMM 优化中 Split-K 技巧的启发。FlashInfer 将 KV 序列维度划分为不同的线程块并将其分派到不同的线程块,并在第二次传递中合并它们。这个相同的想法也由 Flash-Decoding 提出,您可以查看他们的精彩 博客文章,了解可视化和更多详细信息。