级联推理:内存带宽高效的共享前缀批处理解码
许多大型语言模型(LLM)推理任务涉及从共享前缀(prompt)生成多个独立的文本,例如Self-Consistency、Tree of Thoughts和Skeleton-of-thought。使用共享前缀服务LLM可能会消耗大量的内存和时间,特别是当前缀较长且请求数量较多时:一个可能的用例是长文档问答(图1),多个用户与聊天机器人互动,使用相同的文档作为prompt。虽然vLLM通过只存储一份共享前缀的副本缓解了内存问题,但它仍然受限于低效率,因为默认的PageAttention实现没有优化对共享prompt的KV-Cache访问。
在这篇博客文章中,我们介绍了级联推理(Cascade Inference),它简单地解耦了共享前缀和唯一后缀的注意力计算,并允许将共享的KV-Cache存储在GPU共享内存(简称SMEM)中,以便多个请求快速访问。我们展示了级联推理可以极大地加速共享前缀的批处理解码操作,与基线vLLM PageAttention实现相比,在H100 SXM 80GB上实现了高达31倍的加速,与没有级联的FlashInfer批处理解码操作符相比,实现了26倍的加速。这些内核已作为PyTorch和C++ API集成到FlashInfer中。
图1. 为多个用户提供文档问答服务的示例,所有请求共享同一本书作为prompt。
背景
GPU的内存层次结构包括全局内存、L2缓存、SMEM/L1缓存和寄存器。全局内存和L2缓存是所有流多处理器(SM)共享的,而SMEM/L1缓存和寄存器是每个SM私有的。访问全局内存和L2缓存的吞吐量远低于访问SMEM/L1缓存和寄存器的吞吐量。因此,最大程度地减少对全局内存和L2缓存的访问对于实现高吞吐量非常重要。
在CUDA程序中,独立的任务被分配到不同的线程块1,每个线程块由一个SM执行。对于pre-Hopper架构2,每个线程块只能访问其本地的SMEM和寄存器。
多查询和单查询CUDA内核的区别
在多查询注意力(用于prefill/append)内核中,多个查询访问KV-Cache的同一区域。多查询注意力内核的常见实现是在单个线程块中处理多个查询,并将KV-Cache加载到共享内存中,并行计算多个查询与KV-Cache之间的注意力。这种方法带宽高效,并且可以通过利用Tensor Cores来最大化TFLOPs/s,但如果不同查询的KV-Cache不同,则不适用。
另一方面,单查询注意力内核(用于解码)假设每个查询都有自己的KV-Cache,批处理无法提高此操作符的运算强度。在这种情况下,在同一个线程块中处理多个查询没有好处,因为重用KV-Cache的机会有限。大多数解码注意力内核的实现是在单个线程块中处理一个查询,以保证并行性,从而使所有SM都得到充分利用。然而,这种方法内存带宽效率不高,因为每个线程块都需要从全局内存(或L2缓存,如果缓存行之前被命中)加载KV-Cache。
分而治之
多查询注意力内核和单查询注意力内核都不完全适合共享前缀批处理解码。然而,多查询注意力非常适合处理查询与共享前缀之间的注意力,而单查询注意力可以处理查询与唯一后缀之间的注意力。我们能否结合这两种方法的优点?
递归注意力
答案是“是”,如果我们能找到一种方法来“合并”相同查询与共享前缀和唯一后缀的注意力计算。幸运的是,FlashAttention已经表明,通过不仅存储局部注意力结果,还存储归一化比例,并在运行时重新归一化局部注意力结果,可以组合局部softmax/注意力结果。我们在这里用简洁的符号表述这个想法:
假设 $s_i$ 是查询与索引 $i$ 处的键之间的 pre-softmax 注意力得分
\[s_i = \mathbf{q}\mathbf{k}^T_i,\]我们可以将定义从单个索引推广到索引集
\[s(I) = \log\left(\sum_{i\in I} \exp(s_i) \right),\]让我们也将值向量 $\mathbf{v}$ 从索引推广到索引集(请注意,$\mathbf{v}$ 和 $s$ 的推广是自洽的:当 $I$ 等于 $\{i\}$ 时,我们有 $s(I) = s_i$ 和 $\mathbf{v}(I) = \mathbf{v}_i$)
\[\mathbf{v}(I) = \sum_{i\in I}\textrm{softmax}(s_i) \mathbf{v}_i = \frac{\sum_{i\in I}\exp\left(s_i\right)\mathbf{v}_i}{\exp(s(I))},\]$\textrm{softmax}$ 函数被限制在索引集 $I$ 上。注意,$\mathbf{v}(\{1,2,\cdots, n\})$ 是整个序列的自注意力输出。查询与索引集 $I$ 的KV之间的注意力状态可以定义为一个元组 $\begin{bmatrix}\mathbf{v}(I) \\ s(I)\end{bmatrix}$,然后我们可以定义一个二元合并操作符 $\oplus$ 来组合两个状态(在实践中,我们会将 $s$ 减去最大值以保证数值稳定性,这里我们为了简化而省略了这个技巧)
\[\begin{bmatrix}\mathbf{v}(I\cup J)\\s(I\cup J)\end{bmatrix}=\begin{bmatrix}\mathbf{v}(I)\\s(I)\end{bmatrix}\oplus\begin{bmatrix}\mathbf{v}(J)\\s(J)\end{bmatrix}=\begin{bmatrix} \frac{\mathbf{v}(I)\exp(s(I)) + \mathbf{v}(J)\exp(s(J))}{\exp(s(I)) + \exp(s(J))} \\ \log(\exp(s(I)) + \exp(s(J))) \end{bmatrix},\]合并操作符可以推广到任意数量的注意力状态输入
\[\begin{bmatrix}\mathbf{v}(\bigcup_{i=1}^{n}I_i) \\ s(\bigcup_{i=1}^{n}I_i) \end{bmatrix} = \bigoplus_{i=1}^{n}\begin{bmatrix}\mathbf{v}(I_i) \\ s(I_i)\end{bmatrix} = \begin{bmatrix} \sum_{i=1}^{n} \textrm{softmax}(s(I_i))\mathbf{v}(I_i) \\ \log(\sum_{i=1}^{n} \exp (s(I_i))) \end{bmatrix}\]上述n元合并操作符与二元合并操作符一致,并且我们可以证明该操作符具有交换律和结合律。通过合并索引子集的注意力状态,有不同的方式可以获得整个序列的注意力状态,最终结果在数学上是等效的。
图2. 合并注意力状态的不同顺序在数学上是等效的。
递归注意力允许我们将注意力计算分解为多个阶段,不同阶段可以分配给不同的计算单元/设备。FlashInfer和Flash-Decoding中的KV序列分区技巧使用了相同的思想来合并来自不同线程块的部分注意力状态。
级联推理:算法
有了合并操作符,我们可以将不同KV子集上的注意力计算分配给不同的内核实现。对于共享前缀批处理解码注意力,我们提出了以下分而治之算法:
- 使用多查询(prefill/append)注意力内核计算查询与共享前缀的KV-Cache之间的注意力状态。
- 使用批处理解码注意力内核计算查询与唯一后缀的KV-Cache之间的注意力状态。
- 使用合并操作符组合两个注意力状态以获得最终的注意力输出。
图3的左侧解释了整个工作流程,不同颜色的矩形在GPU中由不同的线程块处理。请注意,对于多查询注意力内核,我们通过SMEM或寄存器访问KV-Cache,而对于解码内核,我们只能通过L2缓存或全局内存访问KV-Cache。级联推理允许我们最大限度地重用公共前缀的内存,从而使注意力计算更加内存高效。
图3. 级联推理的工作流程,吞吐量值改编自博客:TPU vs GPU vs Cerebras vs Graphcore: A Fair Comparison between ML Hardware
我们将这种用于共享前缀注意力的分而治之方法称为“级联推理”(Cascade Inference)。
评估
我们在H100 SXM 80GB和A100 PCIE 80GB GPU上评估级联推理。输入形状改编自LLaMA2-7B(32个头,每个头128维)。我们改变三个参数:请求数(批大小)、共享前缀长度和每个请求的唯一后缀长度。基线实现是vLLM 0.2.6中实现的PageAttention内核,我们还展示了没有级联的FlashInfer批处理解码操作符的性能。所有实现的页面大小(或块大小)都固定为16(FlashInfer有/无级联,vLLM PageAttention)。
图4. 在H100 SXM 80GB上相对于vLLM PageAttention的加速比
图5. 在A100 PCIe 80GB上相对于vLLM PageAttention的加速比
图4和图5显示了在级联和非级联设置下FlashInfer内核相对于vLLM实现的归一化性能。FlashInfer内核在这两种设置下都优于vLLM内核,并且在大多数情况下,级联内核相对于非级联推理内核有显著加速。级联推理的好处随着共享前缀长度和批大小的增加而增加(prefill内核主导执行时间),并随着唯一后缀长度的增加而减少(批处理解码内核主导执行时间)。对于非常长的共享prompt(32768),当批大小较大(≥128)且唯一kv-长度较短(≤256)时,解码内核在H100 SXM 80GB上可以获得高达31倍的加速。
评论和未来工作
级联推理的思想可以推广到多个级别(我们在本博客文章中只展示了两个级别)和多个共享前缀。多级别、多共享前缀的级联推理已集成到MLC-Serving中:基于MLC-LLM的通用服务框架,我们将在未来的博客文章中展示端到端加速。
最近,SGLang(一种用于编程LLM的领域特定语言)提出了RadixAttention,其中KV-Cache被组织成一个基数树结构,并且注意力可以通过多级级联推理进一步加速。我们正在与SGLang团队合作,以实现此功能。
引用
@misc{cascade-inference,
title = {Cascade Inference: Memory Bandwidth Efficient Shared Prefix Batch Decoding},
url = {https://flashinfer.cn/2024/02/02/cascade-inference.html},
author = {Ye, Zihao and Lai, Ruihang and Lu, Bo-Ru and Lin, Chien-Yu and Zheng, Size and Chen, Lequn and Chen, Tianqi and Ceze, Luis},
month = {February},
year = {2024}
}

评论