开源项目

FlashKDA

高性能Kimi Delta Attention的CUDA内核实现,基于CUTLASS构建,专为SM90+ GPU优化。由MoonshotAI开源,可直接作为flash-linear-attention的后端使用,支持可变长批处理、门控机制和状态传递,显著提升KDA算子效率。适合需要高效推理/训练Kimi Delta Attention的开发者。

README

FlashKDA

FlashKDA: Flash Kimi Delta Attention — 基于 CUTLASS 构建的高性能 KDA 内核

新闻

  • 2026-04-22 — 深度博客:FlashKDA v1 背后的设计决策,阅读请点这里。

依赖要求

  • SM90 及以上
  • CUDA 12.9 及以上
  • PyTorch 2.4 及以上

安装

git clone https://github.com/MoonshotAI/FlashKDA.git flash-kda
cd flash-kda
git submodule update --init --recursive
pip install -v --no-build-isolation .

默认情况下,构建过程会检测当前 CUDA 设备并只为该架构编译。对于 wheel 或 CI 构建,请显式编译所有支持的架构:

FLASH_KDA_CUDA_ARCHS=all pip install -v --no-build-isolation .

支持的取值包括 auto(默认)、all,或以逗号分隔的架构列表,如 90a,100a。

将 FlashKDA 用作 FLA 后端

安装后,FlashKDA 会被 flash-linear-attention 的 chunk_kda 自动调用。集成详情请参见 fla-org/flash-linear-attention#852。

依赖要求

  1. 安装 flash-linear-attention >= 0.5.0:
    pip install -U flash-linear-attention
    
  2. 在 torch.inference_mode() 上下文中调用 chunk_kda
    import torch
    from fla.ops.kda import chunk_kda
    
    with torch.inference_mode():
        out, final_state = chunk_kda(
            q=q, k=k, v=v, g=g, beta=beta,
            scale=scale,
            initial_state=h0,
            output_final_state=True,
            use_gate_in_kernel=True,
            use_qk_l2norm_in_kernel=True,
            use_beta_sigmoid_in_kernel=True,
            safe_gate=True,
            A_log=A_log, dt_bias=dt_bias,
            lower_bound=lower_bound,
            transpose_state_layout=True,
            cu_seqlens=cu_seqlens,
        )
    

退出: 设置 FLA_FLASH_KDA=0 以回退到 Triton 路径。

调试调度: 添加 logging.basicConfig(level=logging.INFO),可在命中时看到 [FLA Backend] kda.chunk_kda -> flashkda,未命中时看到 ... rejected: <reason>。

性能

参见 BENCHMARK_H20.md。

测试

bash tests/test.sh
  • tests/test_fwd.py — 正确性测试(与 torch 参考实现精确匹配;并与 flash-linear-attention 对比)

内核 API

flash_kda.fwd

flash_kda.fwd(q, k, v, g, beta, scale, out, A_log, dt_bias, lower_bound,
              initial_state=None, final_state=None, cu_seqlens=None)

参数:

参数 数据类型 形状 描述
q bf16 [B, T, H, K] Query
k bf16 [B, T, H, K] Key
v bf16 [B, T, H, V] Value
g bf16 [B, T, H, K] Gate(激活前)
beta bf16 [B, T, H] Beta logits(预激活;内部应用 sigmoid)
scale float scalar 缩放因子
out bf16 [B, T, H, V] 输出张量
A_log fp32 [H] Log-gate 参数
dt_bias fp32 [H, K] Gate 偏置
lower_bound float scalar Gate 下限(范围 -5.0 到 0)
initial_state bf16/fp32/None [B, H, V, K] 或 [N, H, V, K] (可选)初始循环状态
final_state bf16/fp32/None [B, H, V, K] 或 [N, H, V, K] (可选,输出)最终循环状态
cu_seqlens int64 [N+1] (可选)变长批处理中的累积序列长度
  • 当前要求 K = V = 128。
  • initial_state / final_state 接受 None(无状态)、bf16 或 fp32 张量。当两者都提供时,数据类型必须一致。
  • 当提供 cu_seqlens 时,B 必须为 1,T 是所有序列的总长度,initial_state / final_state 的形状为 [N, H, V, K]。
  • 当 cu_seqlens 为 None 时,每个批次元素被视为独立的序列,状态形状为 [B, H, V, K]。

开发

要为 CUDA/C++ 源码设置 IntelliSense(clangd),请运行:

bash setup_clangd.sh

这会生成一个包含正确仓库路径的 .clangd 文件,并将全局 clangd config.yaml 安装到 ~/.config/clangd/。

引用

@misc{flashkda2026,
      title={FlashKDA: Flash Kimi Delta Attention},
      author={Yutian Chen, Zhiyuan Li, Yucheng Wang, Ming Wei},
      year={2026},
      publisher = {GitHub},
      howpublished = {\url{https://github.com/MoonshotAI/FlashKDA}},
}
开源项目MoonshotAI2026-07-29原文

相关内容