FlashKDA:基于 CUTLASS 的高性能 KDA 内核项目

FlashKDA: high-performance Kimi Delta Attention kernels

分支2Tags0
当前项目代码仓暂无内容

FlashKDA

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

新闻

  • 2026-04-22 — 深度解析博客:FlashKDA v1背后的设计决策,阅读地址 here

环境要求

  • 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-attentionchunk_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] 查询
k bf16 [B, T, H, K]
v bf16 [B, T, H, V]
g bf16 [B, T, H, K] 激活前的门控
beta bf16 [B, T, H] Beta 对数(激活前;内部应用 sigmoid)
scale float 标量 缩放因子
out bf16 [B, T, H, V] 输出张量
A_log fp32 [H] 对数门控参数
dt_bias fp32 [H, K] 门控偏置
lower_bound float 标量 门控下界(范围从 -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_seqlensNone 时,每个批处理元素被视为独立序列,状态形状为 [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}},
}

项目介绍

FlashKDA:高性能 Kimi Delta Attention 核心【此简介由AI生成】

定制我的领域
61.24 K121访问 GitHub