ai官小西

FlashKDA:月之暗面开源的 Kimi Delta Attention 高性能内核

FlashKDA:月之暗面开源的 Kimi Delta Attention 高性能内核

长上下文 LLM 的注意力计算是个实打实的工程噩梦。标准自注意力是 O(n²) 的,序列到 128K token 时显存先爆,算力再爆。线性注意力(Linear Attention)把复杂度压到 O(n),但实现层的效率远不如 FlashAttention 成熟——大部分线性注意力实现跑在 Triton 上,与手工优化的 CUDA 内核有数倍的差距。

2026 年 4 月 20 日,月之暗面(MoonshotAI)在 GitHub 上开源了 FlashKDA——专为 Kimi 模型的 Delta Attention(KDA,Kimi Delta Attention)设计的高性能 CUDA 内核,基于 NVIDIA CUTLASS 构建。截至 2026 年 7 月底,项目已获 1,061 Star、101 Fork,MIT 许可证。

这不是又一个通用注意力库。FlashKDA 是 Kimi 生产系统中实际运行的注意力后端——开源意味着你可以直接看到月之暗面是如何解决长上下文推理中的注意力瓶颈的。


一、什么是 Delta Attention?

要理解 FlashKDA,先要理解 Kimi 的选择:为什么是 Delta Attention,而不是标准 Softmax Attention 或其他线性注意力变体?

标准注意力的困境

传统 Transformer 中,每个 token 都要与所有前置 token 计算注意力分数:

Attention(Q, K, V) = softmax(QK^T / √d) · V

这个 O(n²) 的计算量和 O(n²) 的 KV-cache 显存,使得 128K+ 上下文成为工程噩梦。FlashAttention 把计算效率推到极致,但本质上没有改变 O(n²) 的渐近复杂度。

Delta Attention:带门控的增量更新

Delta Attention 是一种线性注意力变体——它维护一个固定大小的循环状态矩阵(recurrent state),每步增量更新,复杂度 O(n·d²)。核心公式可以简化为:

S_t = exp(g_t) · S_{t-1} + k_t^T · v_t
o_t = q_t · S_t · β_t

其中:

  • g_t(gate):控制历史信息的衰减速率,让模型学会"遗忘"
  • β_t(beta):sigmoid 门控的输出权重
  • S_t:形状为 [K, V] 的循环状态矩阵(K=V=128 时,这是 128×128=16K 个元素)

与标准 Softmax Attention 的关键区别:

维度 Softmax Attention Delta Attention
计算复杂度 O(n²·d) O(n·d²)
KV-cache O(n·d),随序列增长 O(d²),固定大小
长程依赖 全局注意力,精确但贵 循环压缩,近似但高效
遗忘机制 无(softmax 归一化隐式处理) 显式门控 g,模型学习衰减

Delta Attention 的思路是:不是每个 token 都需要精确回顾 100K token 前的内容。对于绝大多数推理任务,一个压缩的循环状态足够。但对于需要精确检索远距离信息的任务(代码库问答、长文档引用),标准注意力仍有优势。


二、FlashKDA 的设计决策

FlashKDA v1 做了几个非显而易见的设计决策。这些决策直接来自月之暗面团队在 Kimi 长上下文推理中的实践经验。以下是他们在 2026 年 4 月 20 日的深度解读博文(docs/20260420-flashkda-v1-deep-dive.md)中公开的关键设计点:

2.1 CHUNK = 16:比 Flash Linear Attention 更小

Flash Linear Attention(FLA)使用 CHUNK = 64。FlashKDA 改用 CHUNK = 16,三个考虑因素:

  1. 数值范围适配 bf16。 当门控下限 lower_bound = -5 时,CHUNK = 16 保证了 exp(cumsum(g)) 的值域落在 bf16 的精度的可表示范围内。这意味着不需要复杂的块内重缩放(intra-chunk rescaling)技巧——大 chunk 中常见的数值不稳定问题自然消失。

  2. 矩阵求逆便宜。 逆一个 16×16 矩阵比逆 64×64 矩阵便宜几个数量级。更重要的是,16×16 的逆可以直接通过 Neumann 级数展开计算,无需任何矩阵分解。

  3. SM80 MMA 完整映射。 所有 CHUNK=16 的运算都干净地映射到 SM80 MMA 指令,不依赖任何架构特化特性,跨 GPU 代际移植性好。

2.2 双内核分拆:K1 + K2

FlashKDA 将完整计算拆成两个内核,按各自天然并行轴组织:

  • K1(token 并行,grid = N × H × num_chunks): 门控激活 → L2 归一化 → 衰减应用 → L / Mqk 构造 → 矩阵求逆
  • K2(仅 head 并行,grid = N × H): 逐块 delta 规则循环 → 输出投影 → 滚动状态累积

早期原型是单内核融合的。但 K1 阶段的 token 并行度远高于 K2 的循环阶段——结果是大量 SM 在等 K2 完成时闲置。分拆后,端到端至少有 15% 加速,且每个阶段可以独立调优。

2.3 bf16 循环状态 + fp16 矩阵求逆

这是两个反直觉的精度决策:

为什么用 bf16 存状态? 标准做法是 fp32 存循环状态。FlashKDA 改用 bf16:将状态矩阵的共享内存占用砍半,同时从关键路径上消除了 fp32 → bf16 的类型转换。月之暗面团队的内部测试表明:只要状态更新本身用 fp32 FMA 指令执行,在更新之间以 bf16 存储不会引入可测量的精度损失。

为什么用 fp16 求逆? 16×16 逆矩阵的元素被证明有界在 [-1, 1] 内(参考苏剑林的博客分析),fp16 的动态范围足够。使用 fp16 避免了 bf16 MMA 需要的 fp32 → bf16 转换,且给了 Neumann 级数展开额外的精度余量。

2.4 更多底层优化

  • Base-2 指数。g_act 阶段将指数底换为 2,使用 ex2.approx.ftz.f32 PTX 指令——消除了底数变换 FMA,且 ex2 的吞吐量高于 exp
  • K1 占用率优化。 通过激进的共享内存复用(非重叠生命周期的 union)和 __launch_bounds__(256, 8),用少量寄存器溢出换取了 SM 上线程块的显著增加。
  • 寄存器文件转置。 K2 阶段使用 MOVM_T 指令在寄存器文件中直接转置操作数,消除了阶段之间所有共享内存的中间往返。

三、性能基准

以下是 2026 年 4 月 22 日生成的 H20 基准和 2026 年 5 月 26 日生成的 GB200 基准。所有测试在 T=8192 序列长度、D=128 维度下进行,warmup=30, iters=200, repeats=5

H20(Hopper)

配置 flash_kda (ms) fla_chunk_kda (ms) 加速比
H=96, 定长 2.62 4.84 1.85×
H=96, 变长 batch=[1300,547,…] 2.34 4.83 2.06×
H=96, 变长 batch=1024×8 2.04 4.67 2.29×
H=64, 定长 1.62 3.17 1.95×
H=64, 变长 batch=1024×8 1.40 3.22 2.31×

GB200(Blackwell)

配置 flash_kda (ms) fla_chunk_kda (ms) 加速比
H=96, 定长 1.01 2.33 2.31×
H=96, 变长 batch=[1300,547,…] 0.86 2.33 2.71×
H=96, 变长 batch=1024×8 0.71 2.31 3.27×
H=64, 定长 0.92 1.58 1.70×
H=64, 变长 batch=1024×8 0.48 1.54 3.21×

两个趋势很明显:

  1. 变长批处理场景加速最大。 这是因为 K1 的 token 并行在变长 batch 下有更高的负载均衡优势。
  2. Blackwell 上的相对加速高于 Hopper。 Blackwell 的更大 L1/共享内存和更高的 warp 调度器吞吐量直接放大了 FlashKDA 的融合和状态复用优势。

四、与同类方案的定位对比

方案 目标架构 注意力类型 实现方式 硬件要求
FlashKDA Delta Attention (线性) Kimi KDA CUTLASS CUDA 内核 SM90+, CUDA 12.9+
FlashAttention-3 Softmax Attention 标准/MQA/GQA CUTLASS CUDA 内核 SM80+ (H100 优化)
Flash Linear Attention (FLA) 多种线性注意力 chunk_kda 等 Triton 内核 较宽兼容
FlexAttention (PyTorch) 自定义注意力 用户定义 Torch 编译器 较宽兼容

FlashKDA 不需要和 FlashAttention "竞争"——它们解决的是不同的问题。FlashAttention 优化的是标准 Softmax Attention 的显存和计算效率,但 O(n²) 的复杂度天花板不变。FlashKDA 针对的是线性注意力变体——KDA——在超长上下文场景下的 CUDA 级优化。它填补的是 FlashAttention 在 128K+ 场景下算力不够、FLA 的 Triton 内核效率不够的中间地带。

评分(多维评估)

维度 评分 说明
单内核效率 9/10 深度 CUTLASS 优化,接近手工 PTX 极限
通用性 4/10 仅支持 KDA,且硬性要求 K=V=128、SM90+
集成度 8/10 自动注册为 FLA 后端,一行代码切换
文档质量 7/10 深度解读博文优秀,但缺少多场景案例
许可证开放性 10/10 MIT 许可证,对商业友好

五、我们的判断

FlashKDA 是长上下文 LLM 推理优化的一个示范性的工程案例。它不试图做一个"通用注意力库",而是聚焦一个具体场景(KDA + 长上下文 Kimi),在这个窄领域做到极致。

但它的局限性也同样明确:硬性要求 K=V=128、仅支持 SM90+ GPU(Hopper 及更新)、不支持训练场景(当前只有 forward)。这意味着它的受众非常窄——基本限于使用 Kimi 架构或类似 Delta Attention 变体的团队。

我们的建议:

  • 如果你在用 Kimi 或 FLA 的 chunk_kda 做推理部署,立刻装 FlashKDA。它的集成是自动的(FLA 后端发现机制),性能提升实在。
  • 如果你在做其他线性注意力变体(Mamba、RWKV、RetNet)的推理优化,可以学习 FlashKDA 的设计模式(chunk 选择策略、双内核分拆、精度折衷),但内核本身不能直接用。
  • 如果你只需要标准 Softmax Attention,继续用 FlashAttention-3 就好。FlashKDA 解决的不是你的问题。

FlashKDA 的开源选择——MIT 许可证 + 深度解读博文——在国产大模型厂商中并不多见。月之暗面把一个生产级别的注意力内核完完整整地摊开给你看,这本身就是一种技术自信。


参考资料

  1. MoonshotAI — FlashKDA GitHub 仓库(1.1K+ Star, MIT 许可证)
  2. MoonshotAI — FlashKDA v1 深度解读博文(2026-04-20)
  3. MoonshotAI — H20 基准 (BENCHMARK_H20.md)(2026-04-22)
  4. MoonshotAI — GB200 基准 (BENCHMARK_GB200.md)(2026-05-26)
  5. FLA-org — Flash Linear Attention PR #852: FlashKDA 集成
  6. 苏剑林 — 《让人惊叹的Johnson-Lindenstrauss引理:应用篇》(fp16 矩阵求逆 bound 分析)
  7. Tri Dao et al. — FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision