MoonshotAI/FlashKDA

GitHub: MoonshotAI/FlashKDA

基于 CUTLASS 构建的高性能 KDA(Kimi Delta Attention)CUDA kernel,为 flash-linear-attention 提供加速后端。

Stars: 501 | Forks: 50

# FlashKDA FlashKDA:Flash Kimi Delta Attention — 基于 CUTLASS 构建的高性能 KDA kernel ## 新闻 - **2026-04-22** — 深度解析博客:FlashKDA v1 背后的设计决策,请阅读[这里](docs/20260420-flashkda-v1-deep-dive.md)。 ## 环境要求 - 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` 自动调度(auto-dispatch)。有关集成细节,请参阅 [fla-org/flash-linear-attention#852](https://github.com/fla-org/flash-linear-attention/pull/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: `。 ## 性能表现 请参阅 [BENCHMARK_H20.md](BENCHMARK_H20.md)。 ## 测试 ``` bash tests/test.sh ``` - `tests/test_fwd.py` — 正确性测试(与 torch 参考实现完全匹配;并与 `flash-linear-attention` 进行了对比) ## Kernel 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 | 标量 | 缩放因子 | | `out` | bf16 | `[B, T, H, V]` | 输出 tensor | | `A_log` | fp32 | `[H]` | Log-gate 参数 | | `dt_bias` | fp32 | `[H, K]` | Gate 偏置 | | `lower_bound` | float | 标量 | 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 tensor。如果同时提供这两者,它们的数据类型必须匹配。 - 当提供 `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}}, } ```
标签:AI, CUDA, Vectored Exception Handling, 凭据扫描, 注意力机制, 算子加速, 自动化代码审查, 逆向工具, 高性能计算