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。
依赖要求
- 安装
flash-linear-attention >= 0.5.0:pip install -U flash-linear-attention - 在
torch.inference_mode()上下文中调用chunk_kdaimport 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}},
}