Generated by DALL-E
由 DALL-E 生成

经过四个月的开发,我们非常激动地宣布 FlashInfer 0.2 正式发布。这一重大更新带来了性能提升、更强的灵活性以及关键的错误修复。本次发布的核心亮点包括:

  • 通过 FlashAttention-3 模板 实现更快的稀疏(Page)注意力机制
  • 针对注意力变体的 JIT 编译
  • 支持 多头潜在注意力 (MLA) 解码

具备块/向量稀疏性的 FlashAttention-3 模板

FlashAttention-3 通过巧妙地重叠 Softmax 和矩阵乘法,为 Hopper GPU 带来了突破性的优化。FlashInfer 0.2 集成了 FA-3 模板,在 Hopper 架构上显著提升了 Prefill 注意力性能。

灵活的块稀疏 (Block-Sparsity) 与向量稀疏 (Vector-Sparsity)

FlashInfer 的突出特点是其高度灵活的块稀疏 FlashAttention 实现,支持 任何块大小配置。我们的 PageAttention 算子被实现为 块稀疏注意力算子 (Block-sparse Attention Kernels),其中 page_size 指定了块的列数。在最细粒度下,FlashInfer 支持 向量稀疏 (Vector-sparsity)1 (page_size=1),从而实现精确的内存管理(已在 sglang 中使用)和高效的 KV-Cache Token 剪枝。

通过利用 CuTeCustomStrideComposedLayout 抽象,我们将向量稀疏性扩展到了 FlashAttention-3。受到 CUTLASS gather/scatter 卷积的启发,这是通过对 Producer 内存加载模块进行简单修改而实现的。

性能基准测试

我们对比了两种注意力实现:page_size=1 的 PageAttention2(使用向量稀疏注意力实现)和变长稠密注意力3,并在 FA-2 (v0.1.*) 和 FA-3 (v0.2) 后端下,针对相同的问题规模进行了测试。基准测试使用 head_dim=128causal=True,在不同的 Batch Size (B) 和序列长度 (L) 下,使用高斯分布初始化的输入 Q/K/V 张量。

Performance comparison between dense/sparse attention on FA2&3 template
在 H100 SXM5 上,使用 CUDA 12.4 编译的稠密/向量稀疏注意力在 FA-2 和 FA-3 模板上的性能对比。y 轴:不同设置,x 轴:达到的 TFLOPs/s

结果: 在相同条件下,向量稀疏注意力达到了稠密注意力吞吐量的 90%。FA-3 后端性能始终优于 FA-2。得益于 FlashInfer 稳定的 API,从 FA-2 升级到 FA-3 无需修改代码——只需安装 FlashInfer 0.2。用于复现这些结果的参考基准测试脚本可在 此处 获取。

用于注意力定制化的 JIT 编译

受到 FlexAttention 的启发,FlashInfer 0.2 引入了可定制的编程接口来编译不同的注意力变体。我们在 CUDA/Cutlass 中设计了一个模块化的注意力模板。用户可以通过在注意力变体类中指定 LogitsTransform/QueryTransform 等函数对象(Functors)来定义自定义注意力变体。该类字符串将特化我们预定义的 Jinja 模板,FlashInfer 使用 PyTorch 的 JIT 加载函数来编译并缓存这些算子。像 FlashSigmoid 这样的新变体只需极少代码即可实现。更多案例请参考我们的 JIT 示例

JIT Compilation in FlashInfer 0.2
左图:FlashInfer 中的 JIT 工作流。右图:编译新的注意力变体。

除了支持新的注意力变体外,在 FlashInfer 中支持 JIT 的其他好处还包括:

  • 减小 Wheel 包体积: 在最近的版本中,由于我们预编译了所有注意力变体的组合,FlashInfer 的二进制大小呈指数级增长。我们不得不减少特化配置以使 Wheel 大小可控,但这损害了算子性能(正如在 #602 中观察到的,FlashInfer v0.1.6 的 Prefill 性能甚至不如 v0.1.1,因为我们将编译时参数移至运行时,这会降低性能)。FlashInfer v0.2 通过仅提前编译 核心 算子子集,而将其余大部分注意力变体留给 JIT 编译,解决了这一问题。
  • 轻量化开发: 对于细微的 CUDA 改动,无需重新安装 FlashInfer,只需以 JIT 模式安装 FlashInfer 即可。

我们通过最小化头文件依赖和利用拆分编译优化了 JIT 编译速度。因此,Llama 模型的所有算子都可以在服务器级 CPU 上于 15 秒 内完成 JIT 编译。更多细节请查看我们的 JIT 预热脚本

融合多头潜在注意力 (MLA) 解码算子

多头潜在注意力 (MLA)Deepseek v2 引入,通过将其投影到低秩矩阵来压缩 KV-Cache。由于缺乏优化的算子,为 MLA 实现高吞吐量极具挑战。FlashInfer 社区近期利用 矩阵吸收 (Matrix Absorption) 技巧 实现了一个融合算子,提升了内存效率。详细解释请参见 #551

MLA
FlashInfer 中的 MLA 解码算子工作流

未来的计划包括利用 Tensor Core 加速 MLA 解码,从而惠及投机采样 (Speculative Decoding)。

兼容变长输入的 CUDAGraph

FlashInfer 0.2 通过准确估计资源上限,修复了 Prefill 注意力在捕获和重放阶段 Query 长度变化时与 CUDAGraph 的不兼容问题。现在可以使用 CUDAGraph 来加速投机采样和分块 Prefill (Chunked-prefill) 工作负载中的 FlashInfer 算子。

torch.compile 兼容性

FlashInfer 0.2 遵循 PyTorch 自定义算子标准,确保了与 torch.compile 的兼容性。

打包与 CI/CD

我们现在提供 每日构建版本 (Nightly Builds),以便用户无需等待稳定版即可测试最新功能。

其他显著改进

FusedAddRMSNorm 修复

修复了 FusedAddRMSNorm 中的数值问题,该问题曾导致某些模型输出异常。

集成 Cutlass SM90 Grouped-GEMM

我们将 Cutlass 3.5 SM90 Grouped-GEMM 集成到了 SegmentGEMM API 中,加速了 LoRA 和 MoE 的服务过程。

支持非连续 KV-Cache

KV-Cache 现在可以使用非连续的存储布局,改善了对 卸载 (Offloading) 的支持。

更快的 plan 函数

plan 函数现在使用非阻塞的主机到设备 (Host-to-Device) 内存传输,提升了性能。在 FlashInfer v0.2 之后,建议传递 主机张量 (Host Tensors) 而非设备张量,以减少 plan 函数中的同步。

KV-Cache 追加优化

通过按元素而非按请求进行并行化,提升了小 Batch Size 下的 KV-Cache 追加吞吐量。新 API get_batch_indices_positions 支持此功能。请注意,我们对该 API 进行了一些破坏性更改以适应不同的并行模式。关于新 API 的用法,请参考 我们的基准测试

标准化 RoPE 接口

我们对 RoPE 接口进行了标准化,使其与其他框架保持一致。FlashInfer 采用了 fp32 sin/cos 计算以避免 数值问题

路线图

我们非常感谢社区的热爱与支持。为了提高透明度,我们发布了 开发路线图,您可以在此提供反馈并影响 FlashInfer 的未来。

社区贡献

自 v0.1.6 以来,贡献者人数从 41 人增加到了 52 人。我们感谢以下开发者的贡献: