FlashKDA: high-performance Kimi Delta Attention kernels
当前项目代码仓暂无内容
以下内容由 AI 翻译,如有问题请 点此提交 issue 反馈
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-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] |
查询 |
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_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}},
}