FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving
FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving
发表时间: 2025-05 · arXiv:2501.01005 (MLSys 2025)
原文: https://arxiv.org/abs/2501.01005
作者/机构:Zihao Ye, Lequn Chen, Ruihang Lai, Wuwei Lin, Yineng Zhang, Stephanie Wang, Tianqi Chen, Baris Kasikci, Vinod Grover, Arvind Krishnamurthy, Luis Ceze / University of Washington, DeepMind, OctoAI 等
速读
一句话结论
提出了一种名为 FlashInfer 的大语言模型推理服务注意力引擎,通过统一的块稀疏存储、动态负载均衡调度和即时编译技术,在多种复杂推理场景下显著降低了延迟并提升了吞吐量。
要解决什么问题
大语言模型推理服务中的注意力机制计算面临两大底层卡点。首先是输入动态性导致的计算负载不均衡:在实际服务中,不仅存在上下文预填充和增量解码的混合,还会出现前缀共享、推测解码等复杂场景,导致同一批次内不同请求的 Query 和 KV Cache 长度差异巨大,常规的注意力算子极易出现 GPU 流式多处理器(SM)闲置和负载失衡。其次是存储格式与计算逻辑的碎片化:为了优化显存,业界引入了 PagedAttention、基数树等多种不连续的 KV Cache 存储格式,同时模型层面又不断涌现出分组查询注意力(GQA)、滑动窗口、自定义掩码等注意力变体。现有的注意力库往往只能针对特定格式或变体进行硬编码优化,无法用一套底层逻辑同时兼顾显存访问效率和算子定制化需求。
怎么做的
核心思路是将各种复杂的 KV Cache 物理存储抽象为统一的块稀疏矩阵,并将注意力计算解耦为“运行时动态调度”与“即时编译(JIT)算子”两部分。首先,针对显存访问碎片化问题,FlashInfer 引入了块大小可调的块压缩稀疏行(BSR)格式。无论是分页内存还是基数树,都被映射为统一的稀疏矩阵。为了进一步提升内存效率,设计了可组合格式:将共享前缀对应的 KV Cache 提取为大分块的密集子矩阵以利用高带宽共享内存,而独有部分保留小分块,从而在不移动底层数据的情况下优化访存。其次,针对注意力变体繁多的问题,提供了一套基于 CUDA 和 CUTLASS 的模板。用户只需定义数据变换函数,JIT 编译器就能将其注入模板,生成高度优化的底层代码。在数据加载阶段,算子会根据 BSR 索引将分散在全局显存中的稀疏块拉取到连续的共享内存中,随后复用密集的张量核心指令进行计算。最后,为了解决变长序列带来的负载不均衡,设计了感知动态性的运行时调度器。调度器在 CPU 端根据当前批次的 Query 和 KV 长度,将过长的 KV 序列切分为多个数据块,并基于代价函数分配给不同的线程块(CTA):
各个 CTA 独立计算出局部注意力状态后,利用注意力组合的结合律,通过算子 $\oplus$ 对局部输出 $\mathbf{O}$ 和对数求和指数 $\mathbf{LSE}$ 进行规约,得到最终结果:
效果如何
实验在 NVIDIA A100 和 H100 GPU 上进行,测试了 8B 和 70B 参数的 Llama 3.1 以及 13B 参数的 Vicuna 模型。在端到端服务基准测试中,对比代表编译器生成路线的强基线 SGLang 结合 Triton 算子,FlashInfer 在不同请求率下将词间延迟降低了 29% 至 69%。在长上下文推理场景下,通过 JIT 融合旋转位置编码与注意力计算,相比于未融合的 FlashAttention 官方算子,端到端延迟降低了 28% 至 30%。在并行生成任务中,基于可组合格式分离共享前缀,相比于 MLC-Engine 的单格式基线,获得了 13% 至 17% 的加速。在内核级别,面对长度均匀分布或偏态分布的动态输入,其吞吐量显著超越了代表手写极致优化的 FlashAttention2 和 FlashAttention3 库。该方法也存在局限与代价:在预填充阶段,将稀疏数据收集到连续共享内存的操作会带来约 10% 的性能损耗;此外,目前在 vLLM 框架中集成时,由于宿主机端的 Python 数组操作开销,在某些精度设置下会出现轻微的性能回退,且当前版本仅支持前向推理,暂不支持训练反向传播。
A1 主要贡献
大型语言模型(LLMs)以注意力机制为核心,随着模型规模的扩大,高效的 GPU 注意力算子对于实现高吞吐量和低延迟的推理变得至关重要。在 LLM 服务中,注意力机制需要从存储历史上下文的 KV 缓存中读取数据,并根据当前查询生成输出。然而,为 LLM 服务构建高性能的注意力算子面临着传统训练环境中未曾遇到的两大挑战:
1. 工作负载模式和输入动态的多样性:LLM 服务涉及从用于上下文处理的 prefill(预填充)计算到服务期间的 batched decoding(批量解码)等多种注意力计算模式。多请求处理带来了前缀重用的机会,推测解码中的树形解码也引入了额外的注意力模式。此外,查询长度和 KV 缓存大小在批次内和随时间均会发生变化,朴素的实现会导致负载不平衡问题,最优调度需要算子动态适应。
2. 现代硬件实现需要定制化注意力算子:在内存端,Paged Attention 和 Radix Trees 等高效存储格式对于管理不断增长的 KV 缓存至关重要;在计算端,必须精心设计特定于硬件的流水线和模板,以充分发挥各 GPU 架构的性能潜力。此外,现代 LLM 中注意力机制日益多样化(如分组查询注意力、专用掩码、自定义注意力分数计算),需要灵活可扩展的实现策略。
为了解决上述工作负载多样性和硬件异构性带来的复杂性,本文提出了 FlashInfer,这是一个基于代码生成的注意力引擎,旨在加速 LLM 中的注意力计算。其主要创新点如下:
* 引入块稀疏(Block-Sparse)和可组合格式(Composable Formats):以此解决 KV 缓存存储的异构性问题,实现高效的内存管理和访问。
* 开发可定制的注意力模板:通过即时编译(JIT)技术,适应各种不同的注意力变体,确保高性能执行。
* 设计动态调度框架:在有效管理输入动态性的同时,保持与要求静态配置的 CUDAGraph 的兼容性,从而最大化硬件利用率。
* 全面的性能提升:在标准 LLM 服务基准测试、长上下文推理(延迟降低 $28-30\%$)和并行生成(加速 $13-17\%$)等多种推理场景中,显著提升了算子级别和端到端的性能。
A2 背景知识
FlashAttention 机制
FlashAttention 【Dao et al., Flashattention: Fast and memory-efficient exact attention with io-awareness + 2022 + NeurIPS + http://papers.nips.cc/paper_files/paper/2022/hash/67d57c32e20fd0a7a302cb81d36e40d5-Abstract-Conference.html】是一种在减少内存使用的同时计算精确注意力的算法。在前向传播中,它利 用 online-softmax 技巧【Milakov & Gimelshein, Online normalizer calculation for softmax + 2018 + CoRR + http://arxiv.org/abs/1805.02867】,使用恒定大小的片上内存动态更新注意力输出,从而避免 在 GPU 全局内存中实例化注意力矩阵。FlashAttention2 和 3 进一步优化了针对 Ampere 和 Hopper 架构 GPU 的循环顺序和流水线设计。FlashAttention 的计算强度为 $O \left( \frac { 1 } { 1 / l _ { q o } + 1 / l _ { k v } } \right)$,其中 $l _ { q o }$ 和 $l _ { k v }$ 分别是查询和键值缓存的长度。在 LLM 服务中,由于查询长度等于(prefill)或小于(decode)KV 缓存长度,计算强度简化为 $O ( l _ { q o } )$。多查询注意力(MQA)和分组查询注意力(GQA)通过分组查询共享 KV 缓存条目来优化缓存大小,组大小记为 $g = \frac { H _ { q o } } { H _ { k v } }$,从而将计算强度提升至 $O ( g \cdot l _ { q o } )$。
注意力组合(Attention Composition)
Block-Parallel Transformer (BPT) 证明,通过保留注意力输出及其缩放因子,可以组合相同查询但不同键/值的注意力输出。设 $\mathbf{q}$ 为查询,$\mathcal{T}$ 为索引集合,通过对注意力分数执行 log-sum-exp 操作定义注意力缩放因子:
$ \mathbf { L S E } ( \mathcal { T } ) = \log \sum _ { i \in \mathcal { T } } \exp ( \mathbf { q } \cdot \mathbf { k } _ { i } ) $
对应的注意力输出 $\mathbf { O } ( \mathcal { T } )$ 为:
$ \mathbf { O } ( { \mathcal { T } } ) = \sum _ { i \in { \mathcal { T } } } { \frac { \exp ( \mathbf { q } \cdot \mathbf { k } _ { i } ) } { \exp ( \mathbf { L S E } ( { \mathcal { T } } ) ) } } \cdot \mathbf { v } _ { i } $
将集合 $\mathcal{T}$ 的注意力状态(Attention State)定义为输出和缩放因子的元组:$\left[ \mathbf { O } ( \mathcal { T } ), \mathbf { L S E } ( \mathcal { T } ) \right]^T$。关键在于,$\mathcal { T } \cup \mathcal { I }$ 的注意力状态可以通过组合 $\mathcal { T }$ 和 $\mathcal { I }$ 的状态得出。引入二元运算符 $\oplus$:
$ \biggl [ \mathbf { O } ( \mathcal { T } \cup \mathcal { I } ), \mathbf { L S E } ( \mathcal { T } \cup \mathcal { I } ) \biggr ]^T = \biggl [ \mathbf { O } ( \mathcal { T } ), \mathbf { L S E } ( \mathcal { T } ) \biggr ]^T \oplus \biggl [ \mathbf { O } ( \mathcal { I } ), \mathbf { L S E } ( \mathcal { I } ) \biggr ]^T = \left[ \frac { \exp ( \mathbf { L S E } ( \mathcal { T } ) ) \mathbf { O } ( \mathcal { T } ) + \exp ( \mathbf { L S E } ( \mathcal { I } ) ) \mathbf { O } ( \mathcal { I } ) } { \exp ( \mathbf { L S E } ( \mathcal { T } ) ) + \exp ( \mathbf { L S E } ( \mathcal { I } ) ) }, \log ( \exp ( \mathbf { L S E } ( \mathcal { T } ) ) + \exp ( \mathbf { L S E } ( \mathcal { I } ) ) ) \right]^T $
由于 $\oplus$ 满足结合律和交换律,多组注意力状态可以按任意顺序组合。RingAttention 和 Flash-Decoding 利用此属性卸载部分注意力计算,从而减少内存使用并提高硬件效率。在 FlashInfer 中,采用注意力状态作为注意力操作的标准输出,并将 $\oplus$ 作为这些状态的标准归约运算符。
块/向量稀疏性(Block/Vector Sparsity)
块压缩稀疏行(BSR)是一种硬件友好的稀疏格式,它将非零元素分组为大小为 $( b _ { r } , b _ { c } )$ 的连续矩阵,而不是无结构稀疏中的随机分散。BSR 提高了寄存器重用效率,并与 GPU 上的硬件矩阵乘法单元表现出更好的兼容性,同时能够跳过空块以减少计算开销。传统上,Tensor Core 指令对最小维度(如 16)有要求,导致大多数块稀疏注意力库将块大小限制为 (128, 128) 的倍数。然而,近期研究表明,通过先将行/列收集到连续的共享内存中,然后对这些连续数据应用密集的 Tensor Core 操作,可以使用更小的块大小(如 GEMM 中矩阵 B 的 $(16, 1)$ 或矩阵 A 的 $(1, 16)$)实现高效利用。FlashInfer 基于这些技术,支持具有任意列大小 $B _ { c }$ 的块,在处理各种稀疏模式时提供了极大的灵活性和效率。
A3 方法细节
KV 缓存存储设计
统一格式的块稀疏矩阵
最近的 KV 缓存存储技术(如 PageAttention 和 RadixAttention)采用非连续内存存储,其最小粒度为 $( H , D )$ 张量的块(或 token),其中 $H$ 为头数,$D$ 为隐藏维度。这些结构优化了内存碎片,提升了重用和缓存命中率。FlashInfer 证明了这些多样化的数据结构可以统一在块稀疏格式下。在 FlashInfer 中,查询和输出矩阵作为不带填充的交错张量(ragged tensors)高效存储,这有助于将不同请求的查询和输出紧凑地打包到一个张量中。最初,键和值使用与查询相同的索引指针维护在交错张量中,随后被合并到 KV 缓存中。KV 缓存采用块稀疏行(BSR)格式,其中块大小由应用需求定义:$B _ { r }$ 对应于查询块大小,$B _ { c }$ 由 KV 缓存管理算法指定。FlashInfer 算子实现支持任意 $( B _ { r } , B _ { c } )$ 的值。
用于内存效率的可组合格式
受到 SparseTIR 的启发,FlashInfer 通过可组合格式提高了注意力计算的效率。单一的块稀疏格式受限于固定的块大小,其内存效率受到块内行数 $( B _ { r } )$ 的限制。虽然较大的 $B _ { r }$ 提高了共享内存和寄存器的重用率,但也增加了碎片化。可组合格式设计允许基于先验知识分解 KV 缓存稀疏矩阵。例如,如果某些请求共享一个前缀,KV 缓存中对应的行和列就会形成一个密集的子矩阵。此时,可以使用具有较大 $B _ { r }$ 的块稀疏矩阵来高效存储这些子矩阵。这种方法不需要在 KV 缓存中移动数据,而是计算稀疏子矩阵的索引和索引指针数组。在较大块大小上进行的注意力计算可以使用高速共享内存和寄存器访问共享的 KV 缓存条目,从而显著提高内存效率。
计算抽象设计
FlashInfer 为 FlashAttention 开发了 CUDA/CUTLASS 模板,专为密集和块稀疏矩阵设计,兼容从 Turing 到 Hopper 的 NVIDIA GPU 架构。对于 Ada(sm89) 及以下架构,使用 FlashAttention2 (FA2) 算法;对于 Hopper,使用 FlashAttention3 (FA3) 算法。
从全局内存到共享内存的数据移动
由于 FlashInfer 支持任意块大小,块可能与 Tensor Core 的形状不对齐,因此需要专门的数据加载方法。FlashInfer 将分块(tiles)从分散的全局内存传输到连续的共享内存中,以供密集的 Tensor Core 操作使用。单个 MMA 指令的 Tensor Core 输入可以来自块稀疏矩阵中的不同块。稀疏 KV 缓存地址使用 BSR 矩阵的索引数组计算,而密集地址则使用行索引的仿射变换。KV 缓存的最后一个维度(头部维度 $d$,通常为 128 或 256)保持连续,以维持符合 GPU 缓存行大小的合并内存访问。FlashInfer 使用 128B 宽度的异步复制指令 LDGSTS 以最大化内存带宽。虽然 Hopper 架构中的 TMA (Tensor Memory Accelerator) 可以进一步加速数据移动,但它不支持非仿射内存访问模式。因此,FlashInfer 仅在 Hopper GPU 上对连续 KV 缓存使用 TMA,在其他不适用 TMA 的情况下则回退到 Ampere 风格的异步复制。传输到共享内存后,稀疏和密集的 FlashAttention 实现趋于一致,仅在数据加载模块上有所不同。
不同 Tile 大小的微内核
为了适应 LLM 应用不断变化的计算强度,FlashInfer 实现了多种尺寸的 FA2 算法。传统的 FA2 使用有限的 Tile 大小(如 (128, 64)),这在 A100 的 prefill 阶段是最佳的,但对于查询长度较短的解码阶段则效率低下。FlashInfer 提供 Tile 大小为 $( 1 , 1 6 , 3 2 , 6 4 , 1 2 8 ) \times ( 3 2 , 6 4 , 1 2 8 )$ 的 FA2 内核,并基于硬件资源和工作负载强度进行启发式选择:
1. 确定每个批次的平均查询长度(对于 GQA,查询长度与头部组维度融合),选择满足或超过该长度的最小查询 Tile 大小。
2. 将寄存器和共享内存约束公式化为 K/V Tile 大小的函数,以最大化 SM 资源占用。
对于查询 Tile 大小为 1 的情况,由于 Tensor Core 指令的最小行数 $m$ 为 16,FlashInfer 使用 CUDA Cores 模板,而对于其他查询 Tile 大小则使用 Tensor Cores。对于 FA3,FlashInfer 提供了 64 倍数的行 Tile 大小,以对齐 Hopper 的 WGMMA 要求。块行大小 $B _ { r }$ 与查询 Tile 大小 $T _ { q }$ 对齐。
用于注意力变体的 JIT 编译器
现代 LLM 模型使用了许多标准注意力算法的变体。由于变体数量激增,为每种变体在 CUDA 库中特化内核是不可持续的。FlashInfer 设计了一个可定制的 CUDA 模板和 JIT 编译器,它接受注意力变体规范作为输入,并生成优化的内核代码。变体规范包含以下仿函数(Functors):
* QueryTransform, KeyTransform, ValueTransform:在注意力计算前对 Q/K/V 张量应用的变换。
* OutputTransform:返回前对注意力输出张量应用的变换。
* LogitsTransform, LogitsMask:Softmax 计算前对 logits 张量应用的变换和掩码。
每个仿函数具有固定的签名。JIT 编译器通过将变体类插入模板来生成 CUDA 代码,随后使用 PyTorch 的 JIT 编译器进行编译并注册为自定义操作符。FlashInfer 还可以选择在规范中是否使用 Softmax,从而支持不使用 Softmax 的变体(如 FlashSigmoid)。
动态感知运行时
负载均衡调度
FlashInfer 的负载均衡调度算法旨在通过将工作负载均匀分配到所有 SM 来最小化 SM 空闲时间。该算法以查询/输出和键/值维度的序列长度信息为输入,并生成工作负载与协作线程阵列(CTAs)之间的映射,以及部分和最终输出的索引映射。调度算法如下:
1. 输入:$\{ l _ { q o } ( i ) , l _ { k v } ( i ) \} _ { i }$,查询 Tile 大小 $T _ { q }$。
2. 定义 Tile $l _ { q } , l _ { k v }$ 的成本函数:$\mathrm { c o s t } ( l _ { q } , l _ { k v } ) = \alpha l _ { q } + \beta l _ { k v }$($\alpha, \beta$ 为超参数)。
3. 计算最大 KV 块大小 $L _ { k v }$:$L _ { k v } \gets \frac { \sum _ { i } \lceil \frac { l _ { q o } ( i ) } { T _ { q } } \rceil \cdot l _ { k v } ( i ) } { \# \mathrm { C T A } }$。
4. 将每个查询 Tile 的 KV 拆分为最大大小为 $L _ { k v }$ 的块,为每个块分配工作索引 $w$,长度为 $l _ { k v } ( w )$。
5. 令 $W = \{ ( w , l _ { k v } ( w ) ) \}$ 并按长度降序排序。
6. 初始化优先队列 $Q = \left\{ \left( c , 0 \right) \right\}$,其中 $c$ 是 CTA 索引。
7. 当 $W \neq \emptyset$ 时,循环执行:弹出 $Q$ 中成本最小的 $c$ 和当前成本;弹出 $W$ 中的 $w$ 和长度;计算新成本并分配块 $w$ 给 CTA $c$;将更新后的 $(c, \text{new\_cost})$ 压入 $Q$。
由于长 KV 被拆分为多个块,注意力内核不直接生成最终输出,而是将部分输出存储在用户提供的 workspace 缓冲区中。最终输出是所有块的部分输出的收缩(contraction)。调度器在 CPU 上每个生成步骤运行一次,计算出的计划信息被异步复制到 GPU 端,并作为持久化注意力/收缩内核的输入。
FlashInfer 保证注意力和收缩阶段都与 CUDAGraphs 兼容。这两个阶段被合并为一个持久化内核,消除了内核间的开销。网格大小在编译后固定,确保每次生成的指针相同。
编程接口
FlashInfer 提供了与 vLLM、MLC-Engine 和 SGLang 等现有 LLM 服务框架无缝集成的编程接口。用户通过提供注意力变体规范、任务信息和分配的 workspace 缓冲区来初始化包装器。内核在初始化时进行 JIT 编译并缓存以供重用。对于可组合格式,FlashInfer 会创建多个具有不同块大小的注意力包装器,并捕获在不同的 CUDAGraphs 中。
在运行时,plan 函数处理序列长度数据以生成负载均衡调度计划(在 CPU 上执行,不被 CUDAGraph 捕获)。run 函数使用查询、键、值和缓存的计划数据执行注意力计算。CUDAGraph 可以捕获对 run 函数的调用,将整个注意力生成步骤编译为单个图。
# Create workspace buffer
workspace = torch.empty(...)
seqlen_info.init()
# Compile: create CUDAGraphs
graphs = []
for task_info in task_infos:
# Init: compile kernels according to spec
attn = AttentionWrapper(attn_spec, task_info, workspace)
g = torch.cuda.CUDAGraph()
# Dummy plan
attn.plan(seqlen_info)
# Capture CUDA graphs
with torch.cuda.graph(g):
for i, layer in enumerate(layers):
...
attn.run(...)
...
graphs.append(g)
# Runtime: select the best CUDAGraph
g = select_graph(graphs)
finished = False
# Text generation loop
while not finished:
seqlen_info.update()
# Plan per generation step
attn.plan(seqlen_info)
# Replay CUDA-Graph
g.replay()
A4 实验环境
- 数据集:ShareGPT 数据集、序列长度在 512 到 2048 之间均匀分布的合成工作负载(Variable)、MT-Bench 数据集。
- 模型架构:Llama 3.1 8B、Llama 3.1 70B、Vicuna-13B。
- 硬件配置:NVIDIA A100 40GB SXM 和 H100 80GB SXM GPU(包含 1xH100 和 4xH100 设置)。
- 软件配置:CUDA 12.4,PyTorch 2.4.0,存储和计算均使用 f16 精度。对比基线包括 SGLang v0.3.4 (结合 Triton v3.0)、最新主分支的 FlashAttention (包含 FA2 和 FA3)、以及 MLC-Engine。
A5 实验结果
1. 端到端 LLM 服务性能
* 实验内容:在 SGLang 框架下对比 FlashInfer 后端与 Triton 后端的性能。使用 Llama 3.1 8B (1xH100) 和 70B (4xH100) 模型,在 ShareGPT 和 Variable 数据集上测量 TTFT(首字延迟)和 ITL(词间延迟)。请求速率被调整以保持 P99 TTFT 低于 200ms。
* 实验结果与分析:相比 Triton 后端,FlashInfer 在所有设置下均表现出一致的加速效果,ITL 减少了 $29\% - 69\%$。这证明了 FlashInfer 在延迟敏感的在线服务环境中的高效性(引自 Fig 7)。
2. 应对输入动态性的内核性能
* 实验内容:在固定批次大小为 16 的情况下,对比 FlashInfer 与 FlashAttention 在三种序列长度分布(常量 1024、均匀分布 512-1024、偏态 Zipf 分布)下的带宽和 FLOPs 利用率。
* 实验结果与分析:得益于负载均衡动态调度器和多样的 Tile 大小选择,FlashInfer 在均匀和偏态序列长度分布中显著优于 FlashAttention 内核。特别是在解码阶段,FlashInfer 避免了 FlashAttention 使用次优 Tile 大小的问题(引自 Fig 8)。
3. 长上下文推理的可定制性
* 实验内容:在 Streaming-LLM 算法中,使用 Vicuna-13B 在 MT-Bench 数据集上测试。FlashInfer 仅需 20 行额外代码即可生成融合 RoPE 的注意力内核。对比了融合内核与未融合内核(FA)的端到端 ITL 和内核带宽利用率。
* 实验结果与分析:通过改变 Streaming-LLM 的最近窗口大小,FlashInfer 融合内核实现了 $28-30\%$ 的端到端延迟降低。在内核级别,FlashInfer 融合 RoPE 内核的带宽利用率比未融合版本高出 $1.6 - 3.7\mathrm{x}$,证实了注意力内核可定制性的重要性(引自 Fig 9)。
4. 并行生成性能
* 实验内容:在 MLC-Engine 中实现可组合格式,测试 Llama 3.1 8B 和 70B 模型在 ShareGPT 数据集上的并行生成(如 OpenAI API 的 n 参数)。固定请求率为 16,并行 token 数 $n$ 在 1 到 64 之间变化。
* 实验结果与分析:对于中等程度的并行生成($4 \leq n \leq 32$),可组合格式在 ITL 和 TTFT 上均产生一致的加速。在 $n=4$ 时达到峰值加速:8B 模型 ITL 降低 $13.73\%$,TTFT 降低 $16.41\%$;70B 模型 ITL 降低 $17.42\%$,TTFT 降低 $22.86\%$。当 $n$ 较小或极大时,收益趋于平缓(引自 Fig 10)。
A6 结论
本文提出了 FlashInfer,一个多功能且高效的 LLM 服务注意力引擎。通过统一的块稀疏存储、提高内存效率的可组合格式、用于定制的 JIT 编译以及处理输入动态性的负载均衡调度器,FlashInfer 在各种推理场景中均展现出卓越的内核级和端到端性能。未来的工作计划探索将更高级别的 DSL 编译为 FlashInfer 中的注意力规范,以及向其他后端的代码生成。
A7 附录
A. 分组查询注意力(GQA)的头组融合
GQA 允许多个查询头共享相同的 KV 头。如果为每个查询头分配不同的 GPU 线程块,在查询长度较短时会导致 KV 缓存重用率低下。FlashInfer 提供了头组融合(head-group fusion)策略:将不同的 KV 头映射到单独的线程块,同时将查询头与查询长度维度融合。通过在线程块映射中合并查询头维度和行维度,组内所有查询头只需进行一次 KV 缓存的共享内存加载,从而在短查询长度下实现更好的内存重用和吞吐量提升。
B. 稀疏收集的开销
FlashInfer 测量了稀疏收集模块的性能开销。在固定查询头和 KV 头均为 32,头部维度为 128 的情况下,对于解码内核,稀疏和密集 KV 缓存之间的性能差距可以忽略不计(在 $1\%$ 以内)。对于预填充(prefill)内核,存在大约 $10\%$ 的性能差距。这是因为 FA3 模板中的密集注意力使用 TMA 指令进行键/值加载,而 Hopper 架构的 TMA 不支持任意行索引的稀疏收集。因此,FA3 上的稀疏收集依赖于 Ampere 风格的异步复制指令和手动指针算术,这消耗了更多寄存器并需要更小的 KV-tile 大小以避免寄存器溢出。
C. 后端的选择
对于 NVIDIA GPU,FlashInfer 基于 CUDA/CUTLASS 构建而不是 Triton,原因在于:
1. CUTLASS 支持高级 NVIDIA GPU 功能(如 warp-specialization 和 TMA 指令),而这些在 Triton 中目前处于实验阶段或不支持。
2. CUDA/CUTLASS 提供了比 Triton 的 Tile 级别抽象更细粒度的寄存器级别控制,简化了将底层优化(如 PTX 内联汇编)直接合并到 JIT 模板中的过程。
D. 内存管理
FlashInfer 管理锁页(pinned)主机缓冲区和设备 workspace 缓冲区以存储调度器元数据和 split-k 部分输出。
* CUDAGraph 兼容的 Workspace 布局:为了满足 CUDAGraph 对固定指针地址的要求,FlashInfer 基于调度器元数据和部分输出的上限估计,为其分配最大所需容量。
* Split-K 直写优化:在负载均衡调度器中,仅对 KV 长度较大的请求应用 KV 拆分。KV 长度短的请求不需要拆分,因此它们可以直接将部分输出写入最终输出缓冲区(绕过设备 workspace 缓冲区),从而节省计算和内存。
* Workspace 缓冲区大小估计:部分输出的上限大小为 $2 \times \# \mathrm { C T A } \times T _ { q } \times H _ { q o } \times ( D + 1 )$。默认情况下,CTA 总数设置为 $k \times \# \mathrm { S M }$,以最大化 CTA 级别的占用率。
E. 注意力与其他操作的重叠
在 FlashInfer 中,用户可以通过 plan 函数提供分配给特定操作的 SM 数量,FlashInfer 的负载均衡调度器将相应地分配 Tile,从而实现与 GEMM 或跨设备通信等操作在不同 CUDA 流中的重叠。
F. FP8–FP16 混合精度注意力
为了降低内存带宽和存储成本,FlashInfer 实现了混合精度注意力内核:查询和输出保留在 $\mathtt { f p 1 6 }$,而 KV 缓存存储在 $\mathtt { f p 8 }$ 中。FlashInfer 利用快速数值数组转换器和片段洗牌器来加速反量化并高效处理位宽不匹配。
G. 额外评估
* 与 FlexAttention 的比较:在 AttentionGym 基准测试中,FlashInfer 在因果注意力、Logits SoftCap、ALiBi Bias 和滑动窗口等变体中,特别是在长序列长度下,一致优于 FlexAttention。
* 共享前缀注意力内核:可组合格式在长前缀(如 32k)和大批次大小(如 64)下显著优于单一格式。
* vLLM 集成评估:使用 fp8 KV 缓存时,FlashInfer 后端将 vLLM 的 ITL 降低了约 $13\%$,但在 bf16 下由于主机端的 Python 开销导致了轻微的性能回退。
* 细粒度块稀疏评估:在 Quest 算法(使用细粒度 KV 缓存稀疏性)中,FlashInfer 采用稀疏行收集策略以利用小块大小的密集 Tensor Core,在长序列长度下实现了比 PyTorch SDPA 和 FlexAttention 高达 $20\mathrm{x}$ 的加速。
A8 补充细节
相关工作(Related Work)讨论
* 注意力优化:FlashAttention 及其后续版本(FA2/FA3、FlashDecoding)通过 online-softmax 和 Split-K 优化了长上下文和解码性能。FlashInfer 扩展了这些模板以支持稀疏注意力,并对可变长序列使用了类似 StreamK 的优化。与需要分离 KV 缓存管理的 RelayAttention、ChunkAttention 等不同,FlashInfer 的可组合格式支持统一的页表管理。
* GPU 上的稀疏优化:先前的研究(如 Blocksparse 库、TC-GNN、Magicube)提出了向量稀疏格式以利用 Tensor Cores。FlashInfer 在此基础上进行了改进,支持 FlashAttention 中任意的块大小 $\left( b _ { r } , b _ { c } \right)$。
* 注意力编译器:FlexAttention 提供了一个用户友好的接口,将其编译为基于 Triton 的块稀疏实现。FlashInfer 扩展了其编程接口以支持查询/键变换,并专注于 LLM 服务的向量稀疏性和负载均衡。由于 Triton 性能在某些用例中仍落后于 CUDA/CUTLASS,FlashInfer 生成 CUDA 代码,并可作为 FlexAttention 前向传播的后端。
* 讨论与局限性:目前,FlashInfer 仅支持注意力计算的前向传播。其设计将计算与 Tile 调度解耦,允许通过运行时调度器实现 FlashDecoding 和 Lean Attention 等多样化平铺策略,能够覆盖大多数注意力函数,包括最近的 MLA 和 Linear Attention 的内部注意力组件。
💬 评论讨论
欢迎在这里分享您的想法和见解!